Iterative Reasoning Preference Optimization

Richard Yuanzhe Pang, Weizhe Yuan, Kyunghyun Cho, He He, Sainbayar Sukhbaatar, Jason Weston

Introduction

Preference optimization has proven to give large gains when aligning pre-trained language models to human requirements compared to supervised fine-tuning alone (Ziegler et al., 2019; Stiennon et al., 2020). Offline methods such as DPO (Rafailov et al., 2023) are becoming more popular for their simplicity and efficiency. Recent results have shown that iterative application of such an offline procedure is beneficial, whereby the updated model is used to construct new preference relations that are more informative, and hence improve results further. These methods include Iterative DPO (Xu et al., 2023; Xiong et al., 2023), Self-Rewarding LLMs (Yuan et al., 2024), SPIN (Chen et al., 2024), and other methods (Rosset et al., 2024). Common to these approaches is that they have been shown to perform well on general instruction tuning tasks, but they either make only moderate gains or even decrease the performance on standard reasoning tasks. While other kinds of iterative training methods have been applied successfully to reasoning, particularly involving iteration of supervised fine-tuning (SFT) such as STaR (Zelikman et al., 2022) , RestEM (Singh et al., 2023), and V-STaR (Hosseini et al., 2024)V-STaR does use preference optimization, but for training a separate verifier model., using preference optimization to train the generative reasoning model is not applied in these methods.

In this work we develop an approach to apply iterative preference optimization to reasoning tasks, with a particular focus on Chain-of-Thought (CoT) reasoning (Wu et al., 2023). On each iteration we sample multiple chain-of-thought reasoning steps and final answers over training prompts, and then construct preference pairs such that pair winners have correct answers and pair losers have wrong answers. We then train a variant of DPO that includes a negative log-likelihood (NLL) loss term for the pair winners, which also proves crucial for performance. Given the newly trained model, we then iterate the procedure by generating new pairs, and training again, starting from the previously trained model. We find that reasoning performance improves over multiple iterations until it eventually saturates.

We show that our approach, termed Iterative Reasoning Preference Optimization (Iterative RPO), outperforms a number of baselines, including SFT or applying standard DPO, as well as other baselines from the literature. We see an improvement from 55.6% of zero-shot performance on GSM8K to 81.6% after our Iterative RPO training (or from 70.7% to 88.7% with majority voting out of 32 samples), from 77.8% to 86.7% on ARC-Challenge (without using the provided ARC Corpus), and from 12.5% to 20.8% on MATH (without using the provided pretraining corpus in MATH). We provide ablations that indicate the components that lead to these improvements. Overall, our method provides a simple recipe that has the potential to improve the reasoning ability of LLMs over a wide range of tasks.

Iterative Reasoning Preference Optimization

Our approach first assumes access to a base, typically pretrained or instruction-tuned, language model, a set of training inputs, and the ability to judge the correctness of the final outputs. Given a training input, the language model is expected to generate (i) a set of reasoning steps (Chain-of-Thought), followed by (ii) a final answer to the given problem. We assume that we have access to a correctness measure for the final answer, and not for the correctness of the reasoning steps used to reach that answer. In our experiments, we thus consider datasets where gold labels are provided for training inputs, and a binary reward is derived by the exact match between these labels and the final answer generations. However, our approach could also be applied to settings with more general reward models.

On each iteration, our method consists of two steps, (i) Chain-of-Thought & Answer Generation and (ii) Preference Optimization, as shown in Figure 1. For the ttht^{\text{th}} iteration, we use the current model MtM_{t} in step (i) to generate new data for training the next iteration’s model Mt+1M_{t+1} in step (ii).

We assume we are given an initial model M0M_{0}, and a training set D={xi,yi}D=\{x_{i},y_{i}\} containing questions xix_{i} and their correct answers yiy_{i} . The model will be trained and updated at each iteration, resulting in models M0,M1,…MTM_{0},M_{1},\dots M_{T}.

Chain-of-Thought & Answer Generation

Given the current model MtM_{t}, we generate NN different responses for every input, where each response consists of CoT reasoning cc followed by a final answer yy:

In the general version of our approach, one then computes the reward rinr_{i}^{n} for each of these responses based on the correctness of their answers, i.e., rin=R(yin,yi)r_{i}^{n}=R(y_{i}^{n},y_{i}). In our experiments this simply corresponds to rin=1r_{i}^{n}=1 if yin=yiy_{i}^{n}=y_{i}, and 0 otherwise; i.e., whether the prediction matches the answer provided in the training dataset. Thus we have constructed a set of generated responses augmented with rewards:

Preference Optimization

In the next step, we first construct a dataset of response pairs DtpairsD_{t}^{\text{pairs}} based on the generations GiG_{i} from the current model MtM_{t}. The paired data is constructed such that chosen (winning) responses have higher rewards than rejected (losing) responses. This data is then used for preference optimization. In general, this can be done by selecting two responses for the same input, such that one has higher reward than the other, and setting the one with higher reward as the winner. In the binary reward case, we can split the generated responses GiG_{i} into two sets based on their rewards:

Next we build a dataset of preference pairs by selecting a winner response (ciw,yiw)(c_{i}^{w},y_{i}^{w}) from GiwG_{i}^{w}, and a loser response (cil,yil)(c_{i}^{l},y_{i}^{l}) from GilG_{i}^{l}. In particular, we simply iterate over GiwG_{i}^{w} and GilG_{i}^{l} simultaneouslyIf the iteration reaches the end of a set, it restarts from the first element. to produce KK pairs {wk,lk}\{w_{k},l_{k}\}, in order to ensure we use as much of the data as possible.

Given the preference pairs, we can now train a new model MθM_{\theta} that will become our next model Mt+1M_{t+1}. The parameters θ\theta are initialized from model MtM_{t}, and updated with a loss function that combines the DPO loss (Rafailov et al., 2023) for learning from the preference pairs, and the negative log-likelihood (NLL) loss for learning over the winning response from each pair. The loss corresponding to each preference pair is as follows:

Here M(x)M(x) denotes the probability of sequence xx under the model MM, and σ\sigma is the sigmoid function. We use the previous iteration’s model MtM_{t} as the reference model in the denominator of the DPO term. Note that the NLL term is normalized by the total sequence length. The hyperparameter α\alpha balances the two loss terms. For brevity we omitted the pair index kk, but we optimize this loss on each of the k∈[1,K]k\in[1,K] pairs generated for every input sample. At the end of this training, we thus obtain our next model Mt+1=MθM_{t+1}=M_{\theta}, which will be then used to build data for the subsequent iteration.

Iterative Training

Our overall procedure trains a series of models M1,…,MTM_{1},\dots,M_{T} where each successive model t+1t+1 uses preference data DtpairsD_{t}^{\text{pairs}} created by the ttht^{\text{th}} model.

In our experiments, we define the models, and the training data they use as follows:

: Base LLM; in our experiments we initialize with a fine-tuned instruction following model.

: Initialized with M0M_{0}, then trained with D0pairsD_{0}^{\text{pairs}} using LDPO+NLL\mathcal{L}_{\text{DPO+NLL}}.

: Initialized with M1M_{1}, then trained with D1pairsD_{1}^{\text{pairs}} using LDPO+NLL\mathcal{L}_{\text{DPO+NLL}}.

: Initialized with M2M_{2}, then trained with D2pairsD_{2}^{\text{pairs}} using LDPO+NLL\mathcal{L}_{\text{DPO+NLL}}.

This approach can be seen as a similar, but simpler, instance of the Self-Rewarding LLM training scheme proposed in Yuan et al. (2024), with three differences. Firstly, on each iteration in Self-Rewarding a new set of prompts is created to explore the input distribution, but in our approach we use the same fixed set of prompts. Secondly, due to this choice our experimental setup does not require a sophisticated reward model to judge the model generations, as we assume the training prompts have provided gold labels which we compare to. These two omitted steps are challenging for reasoning tasks because they require a language model to verify correctness, which is known to be difficult (Huang et al., 2023). Thirdly, we show that our DPO+NLL objective is important for our reasoning tasks, whereas Self-Rewarding LLM’s used the standard DPO objective.

Our approach is also related to the iterative training in the Self-Taught Reasoning (STaR) method (Zelikman et al., 2022), except that their approach uses SFT training, rather than preference optimization using DPO-like training. Preference optimization allows the use of negative examples of reasoning chains and answers, which we show improves performance. See section 4 for more discussion of related work.

Experiments

In our first set of experiments, we use the GSM8K dataset (Cobbe et al., 2021) that contains real grade-school math word problems. Each problem contains a question xix_{i}, gold chain-of-thought solution cic_{i}, and a final numerical answer yiy_{i}. For our entire training process, we only use the training set of around 7.5k problems without any extra data.

As a seed model M0M_{0} we use the chat version of Llama-2 70B model (Touvron et al., 2023), which is instruction finetuned. We use a zero-shot prompt containing the question together with instructions to produce a chain-of-thought and to follow a specific format so the final answer can be easily extracted (the exact prompt is given in subsection A.1). In each iteration, we generate N=30N=30 solutions per problem using sampling with temperature 0.8 for iterations 1–2 and temperature 1.3 for iterations 3–4 (hoping that there is a significant number of incorrect generations in later iterations). Since some problems might not have any model-generated correct solution, we include the gold human written solution (ci,yi)(c_{i},y_{i}) in the winning set GiwG_{i}^{w} so it is not empty. Then we generate K=10K=10 pairs per problem for training with our loss in Equation 1, and filtered out examples that were too long in terms of overflowing the context length or else do not have any incorrect generations. This gave around 55–60k pairs for training, per iteration.

In total, we performed 4 iterations, producing models M1M_{1}, M2M_{2}, M3M_{3} and M4M_{4}. For each iteration, we train a maximum of 5000 steps, then select the best checkpoint using a held-out 1k samples from the training set. We then retrain including those 1k samples for the selected number of steps. The coefficient α\alpha is tuned in {0.5, 1, 2, 4} when training M1, and we end up using 1 for all experiments in the paper. We used a batch size of 16 and a learning rate 7e-7.

Overall results are given in Table 1, where we give the exact match accuracy on the GSM8K test set.

We find that Iterative RPO outperforms zero-shot CoT, supervised finetuning (SFT) on the gold CoT solutions and variants of DPO by a wide margin. SFT gives a boost in performance compared to zero-shot CoT from 55.6% to 63.5%, but this is still far from the 81.6% of Iterative RPO. We apply standard DPO to the same set of preference pairs D0pairsD_{0}^{\text{pairs}} as used in the first iteration of our method. Whether initializing from Llama-2-70b-chat (M0M_{0}) or from SFT training on the chosen (winner) examples, we find that DPO performance, while being better than zero-shot CoT, is no better than the SFT model, with accuracies of 61.8% or 60.3% respectively. We also show that SFT on only the chosen CoT solutions, which corresponds to the first iteration of the STaR method, improves results to 65.2% over SFT on the gold solutions alone, but still falls short of the performance of the first iteration of Iterative RPO. One hypothesis for these improvements is the necessity of including the rejected sequences in the training objective, otherwise their probability increases along with the chosen samples, see Figure 3. We note this observation has also been reported in concurrent work (Hong et al., 2024). All of the results reported above are using a single generation at test time using greedy decoding. If we use majority voting over 32 samples (sampling with temperature 0.8), a standard approach to improve performance in the literature, we can improve the accuracy of our approach from 81.1% to 88.2% for iteration 3, and from 81.6% to 88.7% for iteration 4 of Iterative RPO.

Iterations of Iterative RPO yield improve reasoning

We observe that Iterative RPO provides improvements over its training iterations, increasing the base model accuracy by 47% (from 55.6% to 81.6%) in total. In contrast, supervised training using the gold CoT only brings about a 14% accuracy boost. We see performance improves across each iteration, from 73.1% to 78.0% to 81.1% to 81.6%. However, the gain decays across the iterations (17.5%, 4.9%, 3.1%, 0.5%), indicating an upper limit on learning across iterations, especially as we are iterating across a fixed number of prompts, i.e., only from the training samples. We also show that it is the iterations of updating the model (i.e., initializing from the previous model) that are helping, not just because there is more data in the form of new pairs generated from the fixed training set. To test this we run the first iteration of Iterative RPO but on twice as much paired data, as well as the STaR method first iteration with twice as much data as well. In both cases performance improves compared to less data, but not as much as performing two iterations. Iterative RPO with twice as much data obtains 74.8% (an improvement over 73.1% using the original dataset size); however, training for two iterations obtains 78.0%. For STaR, training on twice as much data obtains 66.9%, compared to 65.2% with the original data, which is still a much lower performance than Iterative RPO.

NLL loss is necessary in our method: DPO with NLL vs. DPO without NLL

The first iteration of our method can be compared to standard DPO training, which uses the same preference data, as reported in Table 1. We see a large performance drop (73.1% vs. 61.8%) using DPO compared to our method after 1 iteration. The gap remains large even when the standard DPO training starts from the superior SFT-tuned model, which it has been argued improves DPO’s performance (Rafailov et al., 2023, 2024). Our results support the need of the NLL loss term in our training, not just using SFT for initialization. To further understand this, we plot the sequence-level log probability over training steps for these methods in Figure 2. We see that for DPO without NLL loss there is a decrease over training for the chosen sequences, whereas for DPO with NLL there is not, which may help explain the improved performance of the latter. We note that related observations have been made elsewhere in various settings (Pal et al., 2024; Xu et al., 2024). Further, we note that whether we initialize with Llama-2-70b-chat or SFT on chosen for Iterative RPO, accuracy results of first iteration training do not seem to deviate (both obtain the same score 73.1%).

Other results in the literature

We can compare our results to others in the literature, even if their experiments are in different settings. Touvron et al. (2023) reports an accuracy of 56.8% for 8-shot Llama-2-70b, which is close to our zero-shot CoT results for Llama-2-70b-chat. In terms of closed-source proprietary language models, some results are superior results to ours, while others are not, for example GPT-4 obtains 92.0% (5-shot chain-of-thought) (Achiam et al., 2023), Claude 2 obtains 88.0% (Anthropic Team, 2023), PaLM 2 obtains 80.7% (Anil et al., 2023), while GPT-3.5 obtains 57.1% (5-shot) (Achiam et al., 2023). We note that the size (number of parameters) and the makeup of the training set of some of these models have not been fully disclosed. For results that use the same size and class model, Llama-2-70b, MetaMath (Yu et al., 2023) reports an accuracy of 82.3%, while WizardMath reports 81.6% (Luo et al., 2023). These last two results use additional augmented training data, whereas our method does not use additional prompts. We note that such approaches should be orthogonal to ours, and both can provide benefits.

2 ARC-Challenge task

To test reasoning capabilities outside of mathematics, we employ ARC (Clark et al., 2018) which covers multiple science subjects. The dataset contains 7.7k multiple-choice questions split into easy and challenge sets. We report results on the ARC-Challenge test set which has 1172 examples. There is no gold chain-of-thought reasoning provided for training examples in this task, which in any case is not required in our method, as we only compute rewards based on the final answer. One consequence is that if there is no model-generated correct solution for a question, then that question is not included in our training. We thus follow the same setup as before to first generate reasoning and then a final answer by the models (see subsection A.1 for prompt) to construct data for iterations of Iterative RPO. We only train on the training set and do not utilize ARC Corpus.

Specifically, in each iteration, we generate N=30N=30 solutions per problem using sampling with temperature 0.8 for iterations 1–2 and temperature 1.3 for iteration 3. We select K=20K=20 pairs of solutions per problem. We end up with around 20k example pairs for iteration 1, 11k example pairs for iteration 2, and 5k example pairs for iteration 3. The decrease in the number of examples is due to the lack of incorrect samples for a number of questions in later iterations. Each iteration is trained on a maximum of 4000 steps. The hyperparameter tuning relies on the provided development set.

We hence perform experiments using a very similar setup to the one previously described for GSM8K. Overall results are given in Table 2. We again find that Iterative RPO provides increased performance across iterations (84.8%, 86.2%, 86.7%) over three iterations. Majority voting using the model in the third iteration (32 samples, temperature 0.8) leads to another small boost (87.9%). These results outperform zero-shot CoT (77.8%), SFT on chosen sequences (79.8%) and standard DPO (83.5%).

Even though we arrive at similar conclusions to the ones from GSM8K, we find these results especially noteworthy due to the multiple-choice nature of the task. As there are typically only four possible answers, the generated data in step (i) of Iterative RPO may provide a CoT and a final answer that is correct by luck (as random guessing is correct 25% of the time). Hence, the nature of the task may introduce a significant amount of noise in the CoT generations used in preference optimization in step (ii). Nevertheless, the method seems robust to this issue and we still observe performance gains.

3 MATH task

We also experiment with more advanced math problems using the MATH (Hendrycks et al., 2021) dataset that is composed of 12,500 competition problems. The test set has 5,000 examples. Similar to the GSM8K dataset, a gold CoT solution is provided for each problem, and the gold answers can be matched uniquely to predicted answers after normalization to compute rewards. We do not use the accompanying pretraining data. For each MATH question, we use a few-shot prompt given in subsection A.1 as the input to the language model. In particular, the prompt includes four fixed in-context examples chosen from the training set. The language model needs these demonstrations so that the final answers can be properly formatted in LaTeX.

In each iteration, we generate N=20N=20 solutions per problem using sampling with temperature 0.8 for iterations 1–2 and temperature 1.0 for iteration 3. We select K=15K=15 pairs of solutions per problem, and after filtering out pairs with overly long generations, for each iteration we randomly select around 75k example pairs. We train a maximum of 5000 steps per iteration; other details are similar to GSM8K setups.

Results are given in Table 2. We again find that Iterative RPO provides increased performance across iterations, from 17.7% to 19.9% to 20.8% over three iterations. These results outperform few-shot CoT (12.5%), SFT on chosen sequences (16.8%) and standard DPO (12.4%). In particular, DPO degrades the performance compared to initialization.

Overall, we find on all three distinct tasks we tried, from simpler to more difficult, similar observations about the performance gains exhibited by our method.

Related Work

Several works have implemented iterative reinforcement learning from human feedback (RLHF) with a human-in-the-loop to provide additional labels to retrain the reward model at each iteration, e.g., via Proximal Policy Optimization (PPO) (Schulman et al., 2017), reporting improvements across iterations (Bai et al., 2022; Touvron et al., 2023). Recently, approaches have been proposed to perform iterative alignment without a human-in-the-loop. Iterative DPO (Xu et al., 2023; Xiong et al., 2023) optimizes preference pairs using DPO (Rafailov et al., 2023) at each iteration, and then constructs new preference pairs for the next iteration by generating them using the updated model, and scoring them using a reward model. Other iterative methods than DPO exist as well, such as the Cringe loss (Adolphs et al., 2023), Pairwise Cringe Loss (Xu et al., 2023) and ReST (Gulcehre et al., 2023).

SPIN (Chen et al., 2024) is an Iterative DPO-like framework that uses human labels as the winning response in a pair, and the last iteration’s generations as the losing response in the pair. The authors note this has the limitation that once the model generations reach human performance, they are bottlenecked. Further, each input prompt is required to have a human-annotated generation. In contrast, our work only requires the final answer, but not the reasoning steps, and crucially uses the model to generate both winning and losing Chain-of-Thoughts. Only modest gains on reasoning tasks are reported in their work.

Self-Rewarding LLMs (Yuan et al., 2024) also use Iterative DPO with the LLM itself used as a reward model to construct pairs for each successive iteration. Both that work, and the work of Rosset et al. (2024) and Snorkel AI Team (2023) which do similar iterations but with external reward models, show significant gains on general instruction following tasks. However, again, only modest gains on reasoning tasks are reported.

Methods Improving Reasoning Ability

While a number of approaches have been developed to curate or distill training data for reasoning tasks (Yu et al., 2023; Toshniwal et al., 2024), in this work we focus on learning algorithms which is an orthogonal axis. Expert Iteration assumes a reward model, and repeatedly uses rejection sampling to filter generations and train on them, which is found to match the sample complexity of PPO (Havrilla et al., 2024). STaR (Zelikman et al., 2022) relies on a similar loop: generate rationales to answer many questions, prompted with a few rationale examples; if the generated answers are wrong, try again to generate a rationale given the correct answer; and then fine-tune on all the rationales that ultimately yielded correct answers; and repeat. ReSTEM (Singh et al., 2023) assumes a ground truth verifier and also fine-tunes on filtered samples in a repeated fashion. All these methods rely on finding high-quality samples for SFT-like training, rather than using DPO-like pairwise preference optimization as in our work.

The V-STaR method (Hosseini et al., 2024) trains a verifier using DPO and uses this to filter the generations of a model trained by SFT, rather than using DPO to train the generator, as we do. MAPO (She et al., 2024) also recently utilizes DPO but for multilingual reasoning tasks, where they translate across languages.

Conclusion

We proposed an iterative training algorithm, Iterative Reasoning Preference Optimization, for improving chain-of-thought-based reasoning task performance in LLMs. In each iteration, we generate multiple responses and build preference pairs based on the correctness of their final answers, and then use a modified DPO loss with an additional NLL term for training. Our method does not require human-in-the-loop or extra training data, and remains simple and efficient to implement. The experimental results show large improvements on GMS8K, MATH, and ARC-Challenge over various baselines using the same base model and training data. These results indicate the effectiveness of our recipe of iterative training in improving the reasoning capabilities of LLMs.

Acknowledgments

We thank colleagues at Meta and NYU for valuable discussion: in particular, Angelica Chen, Jing Xu, Abulhair Saparov, Vishakh Padmakumar, Nicholas Lourie, and Nitish Joshi.

References

Appendix A More Details on Experimental Setup

For each GSM8K question, we use the following prompt as the input to the language model:

Your task is to answer the question below. Give step by step reasoning before you answer, and when you’re ready to answer, please use the format "Final answer: …"

MATH.

For each MATH question, we use the following prompt as the input to the language model. In particular, the prompt includes four fixed in-context examples chosen from the training set of MATH. The language model needs these demonstrations so that the final answers can be properly formatted in LaTeX.

Your task is to answer the last question below. Give step by step reasoning before you answer, and when you’re ready to answer, please wrap your answer in \boxed, and conclude using the format "Final answer: …"

Question: [question for the first example]

Solution: [solution for the first example]

Final answer: [answer (e.g., number, formula) here]

Question: [question for the second example]

Solution: [solution for the second example]

Question: [question for the third example]

Solution: [solution for the third example]

Question: [question for the fourth example]

Solution: [solution for the fourth example]

ARC.

For each ARC question, we use the following prompt as the input to the language model, assuming the question has four options (each question has three to five options).

Your task is to answer the question below. Give step by step reasoning before you answer, and when you’re ready to answer, conclude using the format "Final answer: (insert letter here)"