GLoRe: When, Where, and How to Improve LLM Reasoning via Global and Local Refinements
Alex Havrilla, Sharath Raparthy, Christoforus Nalmpantis, Jane Dwivedi-Yu, Maksym Zhuravinskyi, Eric Hambro, Roberta Raileanu
Introduction
State-of-the-art large language models (LLMs) exhibit a wide range of downstream capabilities after pre-training. This includes the ability to refine their reasoning on math, science, or coding problems (OpenAI, 2023; Touvron et al., 2023; Chowdhery et al., 2022). However, under close inspection, this refinement ability is quite brittle, often unable to even identify when a solution needs refinement (Huang et al., 2023). When LLMs do produce successful refinements on hard reasoning tasks this is often due to the incorporation of external forms of feedback, e.g. feedback from humans or code, stronger models, or other tools (Zhou et al., 2023; Gou et al., 2023). In this work, we carefully examine and improve the self-refinement abilities of LLMs on reasoning tasks without any external feedback other than the ground truth answers of the training problems. Notably, this means we make no use of data or feedback from humans or stronger models. To do so we start by heuristically decomposing the refinement problem into three parts: firstly deciding when to refine, then where to refine, and finally how to refine.
Outcome Based Reward Models (ORMs) (Cobbe et al., 2021), first introduced as an estimator of final answer correctness given a question to do solution reranking, are a natural choice for addressing step one. For deciding where to refine, we carefully examine the generalization of ORMs to intermediate steps. We find the accuracy of the underlying data generating policy directly affects the ORM’s ability to learn correctness of intermediate solutions steps. This leads to the ORM often under-estimating the solvability of a problem from an intermediate step . The result is high false-negative rates when used to classify steps with errors. Process Based Reward Models (PRMs) instead are trained to directly estimate the correctness of each step. Yet this requires extensive human labeling of model-generated solution steps as valid or invalid. In an effort to improve our ability to give intermediate step feedback, we introduce the Stepwise ORMs (SORMs) which explicitly predict labels at each step indicating the presence of an error. We generate SORM training data by sampling a student policy many times at a step in solution , labeling as valid if we successfully reach the final answer. From an RL perspective, this can be interpreted as learning (a lower bound of) the optimal value function of the reasoning task via approximation of the optimal policy with rejection sampling. The resulting SORM gives better intermediate step-level feedback, allowing us to give information to the refinement model about both when and where to refine. The refinement model must then only decide how to refine.
We initially train global refinement models capable of refining the entire reasoning trace without any feedback beyond an initial draft solution . The training data is generated synthetically, by pairing correct solutions with incorrect solutions as in Welleck et al. (2022). An evaluation of the global refinement model confirms its inability to correctly identify when to refine, demonstrating the need for an ORM. Reusing the SORM training data, we train a local refinement model which uses the feedback given by the SORM to identify the first incorrect reasoning step. We then compare the performance of global versus local refinements on a test set of incorrect solution drafts, finding similar refinement accuracy but on largely disjoint sets of problems. In this sense the global and local refinement models are complementary, with local refinements often able to solve problems global refinements cannot and vice versa. To obtain our best results we combine both global and local refinements, using the ORM to choose the most promising one by acting as a reranker of both plus the initial draft. Using this strategy, we can improve the accuracy of an already strong RL fine-tuned Llama-2 13B mode from 53% to 65% when greedily sampled.
In summary we make the following contributions:
Decompose the refinement problem into three parts, namely deciding when, where, and how to refine a solution by leveraging reward models (RMs).
Highlight the limitations of ORMs in judging the correctness of intermediate steps, despite their ability to judge the correctness of the final answer.
Introduce the step-wise ORM (SORM) to refine which is trained only on synthetic data and can more accurately evaluate intermediate steps than the ORM.
Propose a new method for refining LLM reasoning that decides when to refine using an ORM, where to refine using a SORM, and how to refine using both global and local refinements. We find the two types of refinement are complementary, each able to solve a large class of problems the other cannot.
Demonstrate performance improvements of up to 12% on GSM8K for a 13B LLaMA-2 model using our approach.
Background
Reasoning: We define a reasoning task as a distribution of (natural language) question/answer pairs . The answer could be either a single final answer, typically a numerical value in case of math problems for ease of evaluation, or include a CoT style solution trace justifying a numerical final answer. We often further write the answer as consisting of atomic steps with the final answer being given on step . The notion of a start of a new "step" is problem dependent but in our case always corresponds to a newline token.
Reward Modeling: Given a reinforcement learning (RL) environment, a reward model can be trained to approximate the reward coming from an action in state (Christiano et al., 2017). In the language setting, reward models are trained to approximate the reward given to a response generated by a LLM (Ouyang et al., 2022). The reward is generally sparse and given at the end of a generation as in the case of RLHF (Christiano et al., 2017; Ziegler et al., 2019) where a contrastive preference model is learned for RL and rejection sampling.
Similar to this is the Outcome-based Reward Model (ORM) first proposed as a final answer verifier used to rerank GSM8K solutions (Cobbe et al., 2021). Formally, we say the ORM estimates where is a question and is a model generated answer. Training data for the ORM is generated by sampling an underlying student model many times on questions from a reasoning task . The ORM is then trained to predict where is prefix of intermediate steps and is any hypothetical continuation of sampled from . i.e., at intermediate steps we may interpret the ORM as estimating the probability of leading to the correct final answer. We may sometimes write to emphasize the ORM’s dependence on its data generating student model . More recently, Process-based Reward Models (PRMs) have been proposed to directly supervise the correctness of each step in a solution (Lightman et al., 2023; Uesato et al., 2022). Formally, we write a PRM predicts where is the last step of .
Refinement: We define a refinement of a draft solution and question as a new solution generated by conditioning on both and . We consider both global refinement models, which take as input only and predict , and local refinement models, which take as input an extra parameter indicating the location of an error in , to predict .
Notation: For the rest of the paper we refer to the pre-trained LLM fine-tuned for downstream tasks as the base model. We fine-tune the base model, either on supervised data or using RL, to produce a student model that generates answers given a question . Sometimes we may also write the student model as a policy implicitly depending on learnable parameters . will be used to denote a dataset for TASK with train split and test split being implicit. We will use to denote a question and to denote solution traces. Sometimes we will write which decomposes the solution trace into intermediate steps . will be used to denote the prefix of steps up to . Additionally we will sometimes use and to represent global and local refinements of . denotes the value function of policy . denotes the optimal value function with dependence on the background task implicit.
Related Works
LLM Reasoning: State-of-the-art (SOTA) large language models (LLMs) (OpenAI, 2023; Touvron et al., 2023; Bai et al., 2022; Chowdhery et al., 2022) demonstrate increasingly impressive abilities on hard reasoning tasks as studied by a wide range of math, science, and code benchmarks (Cobbe et al., 2021; Hendrycks et al., 2021b; Sawada et al., 2023; Liang et al., 2022; Srivastava et al., 2022; Rein et al., 2023; Mialon et al., 2023; Chollet, 2019; Hendrycks et al., 2021a; Austin et al., 2021; Mishra et al., 2022; Patel et al., 2021; Gao et al., 2021). Chain of thought (CoT) (Wei et al., 2022) and related techniques (Chen et al., 2022; Yao et al., 2023a; Besta et al., 2023) have emerged as dominant methods significantly boosting LLM performance on these types of tasks. CoT methods allow LLMs to defer giving their final answer by first generating a "chain of thought" involving intermediate computations needed to correctly solve the problem.
LLM Refinement: Intimately related to reasoning ability is a model’s ability to refine previous answers. This work studies the ability of large language models to self-refine their CoT solutions to math reasoning tasks. Several works (Yao et al., 2022; Madaan et al., 2023; Zhou et al., 2023) demonstrate SOTA LLM self-refining and self-critiquing abilities on a range of tasks via prompting and/or tool usage. However, recent work (Huang et al., 2023) argues even for the strongest models such techniques struggle on hard, open-ended reasoning tasks where the model itself must decide when to stop refinement.
Other papers use hand-crafted data augmentation (Paul et al., 2023) or gather human data (Wang et al., 2023b; Chen, 2023; Lee et al., 2023; Saunders et al., 2022; Schick et al., 2022) while still others use techniques from reinforcement learning to generate critiques (Akyurek et al., 2023; Yao et al., 2023b) for larger models. Most related to us is (Welleck et al., 2022) which trains global refinement models in an implicit reinforcement learning like manner by pairing low-value rollouts with high-value rollouts.
Process-based reward modeling (PRMs) (Uesato et al., 2022; Lightman et al., 2023) gives a denser, step-by-step reward for the "correctness" of a particular step without explicitly modeling the step’s impact on the correctness of the final answer. Both ORMs and PRMs are most often used as rerankers over large numbers of candidate solutions, with PRMs generally outperforming ORMs (Lightman et al., 2023). However, PRMs areexpensive to train, requiring extensive human annotation of each step. Uesato et al. (2022) directly compares the performance of a 70B ORM vs PRM on GSM8K, finding both performing similarly when used as a reward for RL and for reranking. They qualitatively note the ORM appears to somewhat generalize to intermediate steps in a manner similar to a PRM but do not quantitatively ablate this observation over multiple models or tasks. Li et al. (2022) attempt to train synthetic stepwise verifiers similar to a PRM which are then used for Monte Carlo Tree Search. Concurrent work (Wang et al., 2023a) proposes training a synthetic process based reward model in a manner similar to our SORM. They then use the RM downstream for RL fine-tuning and rejection sampling.
In contrast to the above works we conduct a careful comparison of ORM/SORM verification abilities at the step level. We then propose to utilize the ORM/SORM for refinement. We accomplish this by generating fully synthetic stepwise labels which allow us to train both the SORM and refinement models.
Method
We start by decomposing the refinement problem into three stages: First, learning when a draft is correct and when it needs refinement. Second, learning where to begin refinement by identifying the first incorrect step. Third, learning how to correct the initial draft. We can naturally address step one by using the ORM which is trained to predict the probability of a draft being correct. This alleviates some of the difficulty, now only requiring the refiner to identify where and when to refine. Additionally, when doing local refinement, we propose using the (S)ORM to localize the position of the first error. This simplifies the task even more, as now the local refiner must only decide how to fix the error and continue from there.
Ideally we might want to use the ORM to identify where a mistake was made by finding the first step such that i.e. is likely to result in the wrong answer. However, because the ORM is acting as a value function for , it tends to hallucinate error steps simply because it expects the data generating student to fail. For example, if almost always fails problems involving division, the ORM will assign low probability of success to a division problem even before the student takes its first step. In these cases we say the ORM is overly pessimistic. This is not ideal when using the ORM to identify the location of mistakes.
Learning a Step-Wise ORM (SORM): Another natural candidate which could be used to identify mistakes at each step is a Process Based Reward Model (PRM) (Lightman et al., 2023). A PRM estimates the probability of correctness of a step , independently of its impact on the final answer. However, this would be expensive, requiring collecting human annotated samples. Instead, we propose to approximate the optimal value function of the reasoning task. corresponds to the value function of the optimal policy which is able to successfully solve the reasoning task from any logically valid intermediate state . Such an optimal value function would have for a solution prefix with no mistakes, and if the prefix already contains a mistake which will result in an incorrect final answer. We call models we train to directly approximate stepwise ORMs or SORMs.
As discussed in Uesato et al. (2022), the ORM possesses some knowledge of intermediate solution correctness, allowing it to approximate a PRM. However, we find in practice this property is dependent on the size of the base model and the difficulty of the task , with ORMs trained on data from larger students and easier tasks giving better approximations to a PRM. When interpreting the ORM as a value function of the data generating student, this makes sense. A larger, more capable student will better approximate the optimal policy , resulting in a better approximation of the ORM to .
Recall, we assume no access to data from humans or better models for fine-tuning. Thus we must generate all training data synthetically for both global and local refinement. Additionally we must generate data for both the ORM and SORM. We divide our proposed training pipeline in three steps. See Figure 1 for a diagram outlining each step.
To produce base checkpoints from which we can generate ORM/SORM training data and initial refinement drafts we fine-tune models using Expert Iteration (EI) (Silver et al., 2017). This is done by sampling the student model times per question and filtering out rollouts with incorrect final answers. De-duplication is then performed on the remaining samples to construct a new finetuning dataset . We then combine this with any available SFT data producing which we use to again fine-tune the pre-trained model. This process is repeated until the maj@1 score of each subsequent fine-tune converges. Note, the fine-tuning dataset used at step is : the union of rollouts generated at the step with previously generated training data (). In the case of GSM8K we first fine-tune each pre-trained model on the given supervised fine-tuning (SFT) data. For SVAMP, which has no CoT SFT data, we 1-shot prompted the pretrained model to generate solutions used to construct an initial EI dataset. We call the resulting model the student model or student policy . For more details of this training process and resulting models see Section B in the appendix.
We generate ORM training data by sampling the RL fine-tuned student policy times per prompt. As usual, we then label each intermediate step as correct if the final answer is correct and incorrect otherwise. To generate training data for our SORM we sample an approximation of the optimal policy at each step in a model generated solution and check correctness of the final answer. We aim to approximate via rejection sampling of our student policy . Concretely, to produce a training label for a step in model generated rollout , we sample the student policy for rollouts starting from the prefix . This produces verifying traces with correct final answers indicated by . We then label as positive if i.e. we can find the correct final answer starting from . In practice we sample rollouts per step, each generating at most 300 tokens. Otherwise we label as negative. We then train the SORM in exactly the same manner as the ORM, predicting the appropriate label after each step in a solution. See Section G for a comparison of the labels assigned by this process to ground truth human labels.
SORM data post-processing To improve our approximation to the optimal policy via rejection sampling we apply several post-processing steps: 1) If a step has a positive label we set for . I.e. all steps before a positive steps are labeled as positive. This accounts for particularly hard problems where the student is able to find the solution with samples from the step but not any prior step , . 2) We enforce a consistency constraint on the verifying rollouts, requiring each intermediate result computed on step of the solution to be used later on. This helps prevent false positives by requiring a verification to make full use of the previous steps it’s verifying. In practice we implement this by checking for each as a string in the suffix after . 3) We balance the number of positive and negative labels at each prefix length in the training dataset. This is crucial, as otherwise there is an imbalance of positive labels towards the start of solutions and negative labels towards the end. This imbalance is easy for SORMs to exploit, leading to models which almost always predict a positive label in the first few steps a negative label towards the end.
As an additional baseline we consider the Balanced ORM which simply balances the number of positives and negatives per question in the ORM training dataset. This is done in an attempt to mitigate the overly pessimisstic behavior of the ORM described earlier.
Our SORM approximation is motivated by observations from concurrent work which shows our student does not need to engage in too much exploration, i.e. sampling, to solve most problems sufficiently in distribution of pretraining data. This suggests rejection sampling to be capable of providing a decent approximation to the optimal policy. Additionally, the deterministic dynamics of the reasoning environment allows us to only sample once from the optimal policy to compute at a prefix . This further reduces our sampling requirements, while also allowing us to conclude that if rejection sampling can solve the problem from a prefix , then will also solve the problem from . Note of course rejection sampling will be weaker than , resulting in the SORM being an under-approximation of .
To train a local refinement model we need a dataset of the form where is a question, is an initial draft, labels the location of the first error in indicating where to refine, and is a refinement with the correct final answer. In pratice, is communicated to the local refinement as a “[BAD]” token prefixing the incorrect step in the draft. Then, at test time, we need a model predicting to localize errors in the draft. Conveniently, we explicitly train the SORM to predict the correctness of each step in . Thus, to produce we infer the SORM on all steps and return the index of the first step with predicted correctness below a threshold . Further, we can construct a refinement training dataset with error annotations using the SORM dataset. Given an incorrect model rollout we can locate step as containing the first error by identifying as the first zero label in the trace. We then pair with a correct verifying trace from the previous (correct) step . This creates a training pair where we label the first error in as . See Figure 2 for an example.
We construct a dataset for global refinement similarly using the ORM training dataset. This is done by pairing incorrect rollouts with correct rollouts for the same question . This constructs a training tuple . To maintain a format similar to local refinement, we put a token at the very start of the incorrect rollout. We combine both refinement datasets to train a model capable of both global and local refinement.
2 Evaluation
We construct a test set for both the ORM/SORM and refinement models by sampling the student model greedily on test questions from the task . For each benchmark this gives us a test set with prompts of the form where is the problem and is an initial draft. For both benchmarks we refer to this as the test set. To generate intermediate step labels we use the same process as used to generate SORM training data. We evalaute the ORM and SORM on this test set by comparing their predictions to these ground truth labels.
To evaluate the global refinement performance we greedily infer the refiner on each sample and compare the resulting refinement to the ground truth. To evaluate the local refinement model we first annotate each pair with the location of its first error using the ORM or SORM. This forms a triplet which we use to greedily sample the local refiner.
For our best results, we propose to sample both a global refinement and a local refinement for a draft and choose the best solution using the ORM reranker. This strategy stems from our observation that global and local refinements each solve complementary, partially non-overlapping subsets of problems the student initially fails on. Thus combining both refinements with the draft significantly expands the set of problems we can solve. Additionally, using the ORM to rerank refinements allows for a cleaner comparison against a best-of-three baseline from the draft-generating student . See Figure 3 for a diagram of the evaluation pipeline.
We also highlight more exploratory work in the appendix. In the main body we consider only process-based local refinement, which relies on locating reasoning errors in a solution trace. One drawback of this approach is its agnosticism to the abilities of the student model doing refinement. Alternatively, we consider value-based refinement which relies on feedback identifying the step in a solution from which the model has the best chance of succeeding. A comparison to process-based refinement is done in appendix Section J. Additionally, in appendix Section C, we compare refinement training using expert iteration to other RL algorithms with various reward schemes.
Results
We evaluate our refinement pipeline on the GSM8K (Cobbe et al., 2021) and SVAMP (Patel et al., 2021) math word problem benchmarks. We fine-tune Llama-2 7B and 13B to produce all downstream models including the ORM, SORM, and refinement models. Note, the evaluation of each model size is self-contained, not utilizing any data or feedback from models of a different size. maj@1 model scores via greedy sampling will be used to evaluate model performance. Hyperparamters for each phase of training are supplied in Section A of the appendix.
SORMs are better than ORMs at evaluating intermediate answers: On GSM8K the SORM improves over the intermediate step accuracy of the ORM by up to 8% from 73% to 81% (See Table 2). This confirms the ORM does a reasonable job estimating intermediate step correctness but can still be improved, particularly for smaller models on a hard tasks like GSM8K. We’ll see this difference in label accuracy also translates into a difference in refinement final accuracy, where it is critical for the ORM/SORM to reliably identify locations of mistakes. In comparison, the balanced ORM underperforms, having comparable intermediate accuracy to the ORM. This is despite qualitiatively appearing to fix the ORM’s over-pessimism, as the balanced ORM assigns roughly 50% chance of success to all questions. We also examine the types of errors models make, finding the SORMs to have a balanced numbers of false positives and negatives when using a 0.5 as the classification threshold.
ORMs better approximate on easier tasks: On SVAMP the ORM has better step accuracy than on GSM8K (see Table 2), particularly the 13B model. As a result the SORM offers less improvement. Most questions in GSM8K are relatively more difficult, requiring at least 4 steps to solve. In contrast, most questions in SVAMP require at most three key steps. This small number of steps likely makes it easier for the ORM to generalize. Additionally, the EI models trained on SVAMP reach on average 15% higher accuracy than the same sized model on GSM8K. This makes the base student model a closer approximation to on SVAMP, making the ORM a closer approximation to .
The importance of a strong data generating student is further highlighted by the difference in accuracies between 7B and 13B models on SVAMP. The 7B student EI model gets an accuracy of 58%, whereas the 13B model gets an accuracy of 70%. Correspondingly, the 13B ORM model performs much better at on intermediate steps than the 7B model. Yet in contrast the 13B ORM on GSM8K performs slightly worse at intermediate steps than 7B. This is perhaps partially explained by the performance of the 13B EI student on GSM8K which only improves 5% over the 7B student.
ORMs are better than SORMs at evaluating final answers: Despite the SORM being generally better at predicting intermediate steps, it is slightly worse at predicting final answer correctness compared to the ORM. This is true for both benchmarks, with the 13B SORM on GSM8K lagging by 5% (See Table 2). However, part of this difference is likely due to statistical biases the ORM is able to exploit, improving final answer accuracy at the cost of over-pessimism. For example, if the problem involves division, the ORM knows the student is likely to fail and immediately predicts a low probability of success. In contrast the SORM is forced to be more optimistic, attempting to carefully examine the correctness of each intermediate step.
Unfortunately, the inaccuracy of the SORM as a final answer predictor also makes it slightly worse as a final answer reranker. For this reason we always use the ORM whenever reranking candidate drafts and refinements. A more detailed comparison of reranking accuracies on GSM8K is done in Figure 4. Note, this comparison is done using ORMs and SORMs derived from a student model trained using only supervised fine-tuning on GSM8K. Rerank accuracies are computed by sampling the student times and scoring each rollout with the ranker. The rollout with the highest score is then chosen as the final answer.
Figure 4 also plots rerank accuracies for SORM models trained on data without additional postproccessing. The best performing SORM uses only consistent verifying rollouts and per-step balanced labels, justifying these as good postprocessing choices.
2 Evaluating global and local refinements
Now, with a better understanding of our SORMs’ capabilities, we can apply them for refinement. Recall that to decide when to accept a refinement we use the ORM as a reranker on the draft and refinement . When performing local refinement we can additionally use both the ORM and SORM to identify the location of the first mistake in . For the ORM we do this by labeling the first step such that where is a threshold hyperparameter. We identify the first error analogously with the SORM. We report results on both GSM8K and SVAMP test sets in Figure 5. Note, we being evaluation without using the ORM as a reranker. This is done to confirm others’ observations that refiners struggle knowing when to refine on their own.
Both global and local refinement models struggle with knowing when to refine: On both benchmarks global and local refinements show little improvement to overall model accuracy. GSM8K 7B global refinements even decreases overall accuracy, with the other models improving by at most 1%. The local refinements improve overall accuracy more, likely due to the presence of the “[BAD]" token indicating the location (and therefore presence) of the first mistake. This underscores the importance of an ORM for choosing when to refine an incorrect draft. We also note that bigger models produce better refinements.
Global and local refinements fix similar percentages of incorrect drafts: To understand how well our refiners perform when refinement is needed we also report results when applying refinement to only incorrect drafts from the test set in Figure 5. In this case both global and local refinements do much better, improving overall accuracy by an average of 10% on GSM8K and 8% on SVAMP. This demonstrates the refiners have learned how to refine, they simply often do not know when.
It is initially somewhat surprising global refinements are able to fix a similar percentage of drafts as local refinements. Local refinements receive extra information from , presumably strictly improving performance over the global refiner. In reality, the provided is noisy as it must be predicted by an imperfect ORM/SORM. We see that even the difference in label accuracy bewteen the ORM and SORM results in a nontrivial difference in refinement accuracy.
Additionally, global refinements have the advantage of optionally restarting a solution from scratch. A local refinement model is trained to reuse the prefix of a solution preceding a “[BAD]” token under the assumption this prefix has no errors. However, even if this prefix has valid reasoning, it may be a low-value solution path for the student. For example, a student who often fails to correctly divide may benefit from starting the problem from scratch in a way that doesn’t require any use of division. global refinements can take advantage of this, whereas local refinements may be commited to valid reasoning with a low chance of successfully completing. See Figure 2 for examples illustrating this point.
Global and local refinements solve partially disjoint, complementary sets of problems: To better understand how global and local refinements compare we examine the overlap between the problems they correctly solve. The last two rows of Table 3 show that, when combined, global and local refinements can fix 41% of incorrect GSM8K drafts from the 13B student. Alone, global refinement and local refinement with the SORM fixes only 28% of problems. Yet, when taking the best of both types of refinement for the same question, we significantly improve performance across all combinations of benchmarks and model sizes. This shows local refinement is able to solve a large set of problems global refinement cannot, and vice versa. Best performance at test time can then be achieved if we have a way of selecting which of the two refinements is appropriate.
Fortunately, we can use the ORM as a reranker for exactly the task of choosing between global and local refinements. Additionally, we can consider the initial draft as a third possible option as a way of deciding if we want to refine at all. Figure 6 shows the results of reranking the draft, global, and local refinement for each question. Since we are effectively sampling three times, we include as a baseline the best of three (Bo3) samples from the EI student. We additionally report overall accuracy if we had a perfect reranker capable of always choosing the correct solution.
Reranking the draft + refinements improves over the draft accuracy by on average 8% across models and benchmarks. When comparing with the Bo3 baseline we still see significant improvements of around 8% on GSM8K. On SVAMP, reranked Bo3 is a much more competitive baseline, itself giving a large improvement over the draft accuracy. An even bigger improvement can be seen when using an oracle reranker, with the 13B refiner improving 11% over even Bo3 on GSM8K.
Conclusion and Future Work
In this paper we study the use of reward models for both identifying when to refine and where to refine LLM reasoning. We found ORM models generalize to some extent to evaluating the accuracy of intermediate steps on easier reasoning tasks but struggle on harder tasks where the training data generating policy is further from . We then propose to approximate the optimal policy via rejection sampling and post-processing, allowing us to generate training labels for intermediate steps used to train SORM models. We find the SORM generalizes better on intermediate test steps than the ORM, but at the cost of final answer accuracy. We then reused the ORM/SORM training data to train a global/local refinement models. We found each type of refinement strategy helped solve a largely unique set of problems, allowing us to combine both via ORM reranking for best performance.
Future work can be classified as either: 1) improving the reliability and verbosity of local error critiques by providing more information on how to refine or 2) augmenting the type of information local refiners use to generate correct solutions. Our study of both ORMs and SORMs reveals large room for improvement when verifying step level reasoning. Allowing verifier models to generate chains of thought appears to offer some benefit (Dhuliawala et al., 2023). Further augmenting verifying CoT with tools (Zhou et al., 2023) allows GPT-4 to effectively solve MATH (Hendrycks et al., 2021a). But it remains unclear how much GPT-4 relies on the tool to solve the problem versus actually uses the tool to augment its own understanding of why a step is wrong.
Another promising direction treats iterative refinement as a form of in-context exploration similar in spirit to ideas from algorithm distillation (Laskin et al., 2022). Here, the aim is to minimize the number of in-context model rollouts needed to figure out how to refine. This also closely relates to work aiming to augment the exploration abilities of SOTA LLMs, a direction we believe is critical to future success. The right iterative local self-refinement strategies might hopefully allow models to access complex behaviors previously inaccessible with naieve iid repeated sampling.
References
Appendix A Hyperparamters
See Table 4 for a list of training hyperparameters used in each training job.
Appendix B RL for Reasoning
In order to start from the best student possible we RL fine-tune the base-model using Expert Iteration. See Table 5 for maj@1 (greedy), maj@96, Rerank@96 and pass@96 scores for EI fine-tuned models.
Appendix C RL for (global) refinement
Setup: We compare the utility of PPO versus EI for refinement on the GSM8K benchmark. To train EI models we sample the SFT2 model trained in Section LABEL:sec:reasoning-rl times per prompt in the train set. We then pair together all incorrect solutions with correct solutions for a fixed question to form training tuples . We then fine-tune from Llama-2 7B to predict with the standard cross-entropy loss for a single epoch. We use an initial learning rate of 5e-5 decaying to 5e-7.
We initialize the PPO model from the SFT2 checkpoint used in Section 2 above and use the same PPO parameters as when fine-tuning from the SFT checkpoint on GSM8K. A single example, included in the appendix for reference, is used to prompt the model for refinement. During training the student model is given a question and draft , where the draft is generated by the SFT model, and tasked with generating a refinement . We give as a sparse reward at the end of the rollout. We additionally experimented with mixing in scores from an ORM, giving a final reward of .
We evaluate all refinement models on a test set with questions from GSM8K test and drafts generated by SFT2. Results are reported in Table 6. Sampling at test time is done greedily, so only maj@1 accuracy is reported. We additionally report a 1-shot prompted Llama-2 7B as a baseline.
All models struggle learning when to refine Our best refinement model is the EI model. However, EI improves over the SFT baseline by only 3%, with PPO showing no improvement at all. This is because both models struggle correctly deciding when to refine. Often the EI model chooses to incorrectly refine correct drafts. The prompted pretrained model does even worse, having not been trained on GSM8K and thus struggling with when to refine and how to produce correct alternative refinements.
The PPO model collapses to simply returning the draft final answer, at least avoiding any negative rewards from incorrectly refining a correct draft. The prompted baseline also exhibits this copying behavior, accounting for the majority of its nontrivial accuracy. We experiment with alternative RL setups for preventing this degenerate behavior by using the ORM as an additional reward, removing the penalty for refining correct drafts, and/or having the student generate both the draft and refinement . However, in all cases the model’s copy bias continues to limit exploration, causing a collapse of the refinement to the initial draft.
Discussion: The results above highlight several failure modes. Firstly, models struggle with determining when to refine, often defaulting to never refining at all. Secondly, when the model does correctly choose where to refine it still struggles with knowing where to refine. Finally, even knowing when and where to refine, the model still must decide how to refine.
In order to improve model refinement capability we propose to decompose the problem by using unique models to solve each failure mode. Fortunately, deciding when to refine can naturally be handled by the ORM which is explicitly trained to predict when a final answer is correct. Additionally, when doing local refinement, we can use the SORM to identify where to refine. This now only requires the refinement model we train to decide how to refine, making the task significantly easier.
Appendix D Misc. Objectives for Reranking
In Lightman et al. (2023) the PRM is used for reranking by estimating for each step and taking the product. Inspired by this, we experimented with a number of different weightings for SORM intermediate step estimates when doing final answer reranking. For a solution of the form these heuristics included:
Mean:
Weighted mean:
Penultimate mean:
The results are plotted in Figure 7. Overall using only the final ORM estimates gives the best reranking accuracy with the penultimate mean coming in at a close second. The weighted mean significantly underperforms all other strategies, even taking the minimum ORM estimate.
Appendix E ORM and SORM extra-model generalization
Both ORMs and SORMs exhibit signs of overfit to the data generating student . When evaluated on GSM8K train set a 7B ORM model incorrectly classifies 42% of correct solutions as incorrect. To examine this more closely we take two base student models, EI (trained with expert iteration) and SFT (supervised fine-tuning), and use both to generate training data for and respectively. we then evaluate both ORMs on test sets generated by each model. Results are reported in Table 7. We find both ORMs underperform on the test dataset generated by the opposite student model.
Appendix F Contrastive vs. Classifier RMs
Both the ORM and SORM are trained as classifiers to predict the probability of a good label at each intermediate step . However, in RLHF there are only preference comparisons over solutions. So the RLHF reward model is often trained via a contrastive loss . We explore the use of a contrastive reward model in the reasoning setting, comparing reranking performance with the ORM. Training data is sampled from a 7B SFT student model with K = 96 rollouts per training prompt at a temperature T = 1.0. We assign ORM step labels in the usual way, setting at a step if and otherwise . To construct preference pairs we select the maximal equal number of positive and negative solutions for the same prompt and form pairs with no solution repeated. Only the contrastive loss on the final token is backpropagated. We then rerank solutions on the test set using scores assigned by both the classifier ORM and contrastive ORM. We find the classifier ORM gets rerank accuracy whereas the contrastive ORM gets , suggesting the classifier to be a strictly better reranker.
Appendix G Accuracy of the SORM data generating method
The SORM data generation process will suffer from both false positives and false negatives. False positives may occur when the student model solves the problem incorrectly but gets the right final answer. False negatives will occur when the rejection sampled student is simply unable to solve the problem from a prefix desipte the prefix being logically valid. In order to verify the correctness of the SORM step-level verification process we hand-label several model generated solutions on GSM8K and compute how well the our ground truth labels align with the generated labels. Over and a total of steps we find the SORM data labels agree with our ground truth 94% of the time.
Appendix H Mixing other sources of PRM data
We additionally experiment with mixing in PRM data on the MATH dataset from Lightman et al. (2023). We train Llama-2 7B to predict negative, neutral and good labels for each step, with “bad” steps being incorrect, “neutral” steps neither making forward nor backward progress, and “good” steps being correct and useful to solving the problem. The resulting PRM gets accuracy on the MATH PRM test set. However, we find the PRM transfers poorly as a final answer correctness predictor, getting only accuracy on an EI generated test set.
Appendix I Self-supervised learning with the SORM
It is likely the SORM dataset generation process is fairly noisy. Low quality samples will directly impact the performance of the downstream SORM, making it critical to improve dataset quality. In an attempt to remove noisy outliers we filtered version of the SORM dataset via SORM self-supervsion. For each training pair , where is the question, is a solution prefix with steps, and is the correctness label, we apply . This generates a self-supervised label . We then filter out all training samples with .
We filter the SORM dataset with a SORM checkpoint trained for 1 epoch and another trained for 2 epochs. The first model, denoted as SORM1, has 75% accuracy on the SORM test set but 91% on the SORM train set. SORM2 gets 78% test but 95% on the train set. It could be SORM2 is overfit to the train set, so we train new SORM models on both filtered datasets. SORM, trained on SORM data filtered with SORM1, gets 79% accuracy. SORM, trained on SORM data filtered with SORM2, gets the same.
Appendix J Value refinement
The local refinement strategy employed in the main body of this work uses critiques attempting to locate steps with logical errors. This can be interpreted as a type of process-based refinement which which gives feedback agnostic to the abilities of the refinement model. More sophisticated forms of feedback might take the underlying capabilities of the model into account, maximizing the chances of this particular student succeeding.
One alternative refinement strategy which gives student specific feedback is value-based refinement. A value-based refinement model receives feedback in the form of a “[BAD]” token at the step in a solution with the highest value for the model. Recall the value of a step is the probability the model gets the correct answer from . Note this is not the same thing as process-based feedback as the step after the highest value step may not necessarily contain an error. Instead, may attempt to solve the problem is a difficult way, for example using division with a model which struggles dividing correctly.
Training a value-based refinement model Further recall the ORM directly estimates the value function of its data generating policy such that for a prefix with . Given a student model on a reasoning task we generate an ORM training set by sampling each prompt in times. We train the ORM as a classifier, setting an intermediate step label where .
To construct the value based refinement dataset we start by reusing the SORM dataset generated as above using policy and rejection sampling. For a sample we identify the highest value step by choosing the step with the most correct verifying rollouts . We then select one of the verifying rollouts whose first step differs from as the improving refinement . This forms a value-refinement training pair where is a “[BAD]” token inserted before step in . We then train a value-based local refinement model by minimizing with the standard cross-entropy loss.
In practice we generate all downstream models and datasets using LLama-2 7B EI on GSM8K as .
Results: To evaluate we use the 7B EI model trained on GSM8K from our best SFT checkpoint. We greedily sample solution drafts on the GSM8K test set, forming a test set for the value-based local refinement model. We then label the highest value step of each draft using the ORM, placing the “[BAD]” token as a prefix to .
We evaluate only on incorrect drafts, comparing directly the performance of process-based SORM refinement. Value-based refinement fixes 14% of incorrect drafts whereas the SORM baseline fixes 21%. Surprisingly, even global refinement outperforms value-based refinement by 6%. This again take this to suggest intermediate ORM estimates are fairly noisy on test data.