Solving math word problems with process- and outcome-based feedback
Jonathan Uesato, Nate Kushman, Ramana Kumar, Francis Song, Noah Siegel, Lisa Wang, Antonia Creswell, Geoffrey Irving, Irina Higgins
Introduction
Recent work has shown that asking language models to use step-by-step reasoning improves performance on reasoning tasks (Shwartz et al., 2020; Nakano et al., 2021; Cobbe et al., 2021; Wei et al., 2022; Kojima et al., 2022; Lewkowycz et al., 2022). While these works have primarily focused on prompting language models, prior work suggests that finetuning should outperform prompting alone (Stiennon et al., 2020; Perez et al., 2021; Ouyang et al., 2022). This raises the question of how best to supervise such models. Two natural approaches are outcome-based approaches, which supervise the final result, and process-based approaches, which supervise each step of the reasoning process, including the last step outputting the final result.
Process-based approaches emphasize human understanding — in order to demonstrate or select good reasoning steps, human annotators need to understand the task. Human-comprehensibility is of direct interest in many domains. For example, in educational settings, an answer without an (understandable) explanation may often confuse more than it explains. In the longer term, human-comprehensibility may also help with detecting when ML systems may be using deceptive or unethical actions to achieve superficially appealing outcomes, such as by subtly manipulating people or systems to increase various metrics (Amodei et al., 2016). Recent work suggests that outcome-based approaches often lack in this area. For example, recent work on natural-language-based reasoning (Zelikman et al., 2022; Creswell et al., 2022) suggests that models optimized exclusively for final-answer correctness can produce the correct final answer, even when their generated reasoning traces are incorrect. Similarly, work on AI safety (Stuhlmüller and Byun, 2022; Krakovna et al., 2020) suggests that such optimization may result in models which execute difficult-to-understand strategies.
This suggests that the choice of supervision approach for language models (LMs) with verbalized reasoning traces likely has important consequences. In this work, we conduct the first comprehensive comparison between process- and outcome-based approaches trained on a natural language task. For this, we use the recently proposed GSM8K dataset (Cobbe et al., 2021) of math word problems. In all cases, we generate a sequence of reasoning steps leading to the final answer, but vary whether or not supervision is provided only on the final answers (outcome-based) or on individual reasoning steps (process-based). While the limited scope of the dataset prevents studying certain safety problems, the dataset enables a clean comparison between outcome- and process-based approaches. For process-based approaches we consider supervision provided by both offline human-generated reasoning traces from the GSM8K dataset itself, as well as online human correctness annotations, which we collect for reasoning steps within model-generated samples.
We compare these approaches in the context of a number of different modeling and training components, including: few-shot prompting, supervised fine-tuning, reinforcement learning (RL) via expert iteration and reward modeling for both reranking and RL. All of our models are based on a large pre-trained LM (Hoffmann et al., 2022).
Throughout, we consider two primary metrics: trace error rate, which measures how often the model makes any mistake in its reasoning trace according to human annotators, and final-answer error rate, which only considers the model’s final answer and ignores the reasoning trace. By “reasoning trace” we refer to all textual steps of reasoning, including the last step which in GSM8K is the final numeric answer.
We find that our best approach, which combines supervised learning with reward-model-based reinforcement learning, significantly improves the state-of-the-art for both trace error rate, from 14.0% 3.4%, and final-answer error rate, from 16.8% 12.7%. Final-answer error rate is further lowered to 2.7% when the model is allowed to abstain on 30% of questions. Our key findings regarding process- and outcome-based feedback are as follows:
Outcome-based and process-based approaches lead to similar final-answer error rates. Both without reward models (23.5% vs. 22.3%) and with reward models (16.6% vs. 14.8%), LMs supervised with final-answer correctness attain nearly the same final-answer error rate as those trained to imitate human-provided solutions.
Both process- and outcome-supervised reward models learn to emulate process-based feedback. Somewhat surprisingly, we find that even reward models trained with outcome-based labels (indicating whether the final answer is correct), result in predictions that agree more closely with the process-based labels (indicating whether each reasoning step is correct) than they do with the outcome-based labels themselves. While this effect may be dataset-specific, as discussed in Section 3, it helps explain the effectiveness of reward models for improving trace error, and we hope that it is investigated further in future work.
Low trace error requires either process-based feedback, or a reward model that emulates it. All models using reinforcement learning directly against final-answer correctness resulted in high trace error, with a best trace error of 12.4%, compared to only 3.8% for our best process-based method. Building on our previous finding, reinforcement learning against a reward model rather than final-answer correctness closes much of this gap, reducing trace error to 5.5%.
In the rest of this paper, we describe the approaches we compare in Section 2 and our results in Section 3. Section 4 discusses implications for process- and outcome-based feedback more generally, Section 5 discusses related work, and Section 6 concludes.
Problem and methods
This section describes the dataset, evaluation metrics and the different modelling components evaluated in this paper. See Fig. 1 for an overview of how they all fit together.
We conduct all experiments on the GSM8K dataset (Cobbe et al., 2021), composed of grade school math word problems. We chose GSM8K because it is a competitive benchmark, and contains natural language reasoning traces. We focus on a single dataset, since the need to recruit human annotators with the domain expertise to accurately evaluate reasoning traces imposes a large up-front cost. Table 2 and Appendix A show several example problems. We split out our own validation set of 256 examples from the original training set, which leaves us with 7118 training and 1319 test examples.
We report two main metrics for all methods evaluated on the GSM8K test set. Final-answer error rate is the fraction of problems for which the method does not produce the correct final answer. Because all final answers on GSM8K are integers, this can be measured with exact string matching. Trace error rate is the fraction of problems with correct final answers for which the method produces at least one incorrect reasoning step. We estimate this via human annotations of the correctness of each reasoning step, using the rating interface discussed in Section 2.7.
We report final-answer and trace errors as two separate metrics because, from a safety perspective, we are particularly interested in errors which remain undetected after applying easy-to-compute proxy metrics (in this case, final-answer errors). For example in an educational setting it is important to show a student the correct steps to get the answer, and we can easily filter out incorrect traces that lead to the wrong answer, but it is much more difficult to filter out incorrect traces that lead to the correct answer. We additionally report Selective final-answer error rate to assess performance when abstaining is allowed and Out-of-distribution (OOD) error rate on pre-algebra problems from the MATH dataset to assess generalization, as described in Sections 3.4 and 3.5, respectively.
2 Training: Overview
Our goal is to train a system for the sequence-to-sequence task (Sutskever et al., 2014) of taking the text of a problem as input and generating the text of an answer as output. For math word problems, the answer is a full reasoning trace: a newline-separated sequence of steps, where the last step is expected to provide the final answer. For GSM8K, the final answer is always an integer.
Our approach broadly follows prior work on RL for LMs (Ziegler et al., 2019; Nakano et al., 2021; Menick et al., 2022). We use an LM as a policy, which maps the problem statement and steps-so-far to a next step. In the RL formalism, this treats each step as an action, and the observation is provided by all the tokens so far. The policy can be obtained through any of few-shot prompting, supervised finetuning (Section 2.3), or RL (Section 2.6). We also train LMs as reward models (Section 2.4), which score proposed completions or partial completions from the policy, and can be used both for reranking samples from the policy, or as the source of rewards during reinforcement learning. In the following subsections, we describe how we train and assemble these components.
3 Supervised finetuning
In supervised finetuning (SFT), we finetune an LM to maximize the log-likelihood of a sequence of target tokens, given a sequence of input tokens. In our paper, we use SFT as a process-based approach by taking the reasoning traces provided in the GSM8K dataset as the target tokens (as opposed to the outcome-based approach of using only the final answer as the target), with the problem statement as the input tokens.
We finetune using AdamW (Loshchilov and Hutter, 2017) with a learning rate of and a batch size of 256. We stop finetuning once the language modeling loss begins to increase on the validation set. For our SFT model, this happens after 70 steps, amounting to slightly more than 2 training set epochs.
4 Reward models
We evaluate two main approaches to training reward models (RMs) (Christiano et al., 2017; Ziegler et al., 2019; Menick et al., 2022), also known as verifiers (Cobbe et al., 2021). In both approaches, we implement the RM as a LM, trained to predict a binary label as either a ‘correct’ or ‘incorrect’ token after each step. In the outcome-supervised RM (ORM), the binary label for each step indicates whether the resulting final answer of that full sample matched the reference final answer, as proposed by Cobbe et al. (2021). A policy which maximizes the ORM score at each step thus maximizes the RM-estimated probability at each step of eventually reaching the correct final answer. For the process-supervised RM (PRM), the binary label after each step indicates whether the steps so far are correct. Because we lack reliable programmatic means for determining the correctness of intermediate steps, we use human annotations for these labels, as described in Section 2.7. A policy which maximizes the PRM score thus selects each step to maximize the RM-estimated probability of the steps so far being correct. If the steps so far are correct, this typically means such a policy minimizes the probability of introducing a mistake on the current step. As reported in Section 3.2, we find this outperforms the approach from Li et al. (2022), which is similar to our PRM but replaces human evaluations with a heuristic based on string matching the results of the intermediate calculations.
Unless otherwise noted, for all approaches which include an ORM, we train the ORM using samples from the policy for that approach, taking samples with temperature . We follow Cobbe et al. (2021) and regularize with dropout, with a dropout parameter of , and otherwise reuse the hyperparameters used for SFT from Section 2.3. To speed up learning in the SFT-based approaches, we initialize the ORM training using the SFT model parameters, while for the few-shot based approaches we initialize from the base pretrained LM. For the PRM, we annotate samples per problem from the SFT policy, restricting to problems where the SFT majority prediction (see Section 2.5) was incorrect, in order to make the most of our human annotation budget. Due to the small size of our human-annotated dataset (1560 full solutions), we initialize the PRM parameters to the ORM parameters and lower the learning rate to . The RM loss curves have some fluctuation, and so we select the RM with the best validation loss before steps.
5 Decoding
For all test-time decoding, we first generate samples of full solutions, and then select the best sample, either by ensembling across samples or by using an RM. In early experiments, we also tried RM reranking after each generated step (rather than the full solution), but found that this led to slightly worse performance, increasing final-answer error by 1-2%. We sample with temperature , and use the syntax from Cobbe et al. (2021) to allow the model to decide when to use a calculator.
We use two approaches to select the best sample. When no RM is available, we use majority voting. For this, we first select the most common final answer from the samples, then select a random sample from among those yielding this selected final answer. This is called self-consistency by Wang et al. (2022), and is similar to more general techniques like Minimum Bayes Risk decoding (Kumar and Byrne, 2004). Otherwise, we use RM-weighted decoding, also called verifier-voting by Li et al. (2022). Here, we weight each sample according to the RM-estimated correctness probability, select the final answer with the largest total weight, and then select the sample with the highest RM score from those yielding the selected final answer. More formally, we select the final answer , where are the model samples, then select the best sample according to . This works slightly better compared to simply selecting the sample with the highest RM score (about 1% final-answer error with the SFT model, slightly more with RL). However, we note that both majority voting and RM-weighted decoding are slightly less general due to their reliance on exact string-matching between final answers.
6 RL via Expert Iteration
All our RL experiments use expert iteration (Silver et al., 2017; Anthony et al., 2017). As a meta-algorithm, expert iteration alternates between two high-level operations. In policy improvement, we combine a base policy with a search procedure to produce samples from a so-called expert policy. Then in distillation, we perform supervised learning on these expert samples to improve the base policy towards the expert policy. We use 5 epochs and select the best model of the 5, based on final-answer test error with RM-weighted decoding, or majority voting if no RM is available.
The initial base policy can be either the SFT policy, or a 5-shot prompted version of our base LM. We particularly note that, aside from the 5 random training examples used for the prompt, none of the few-shot-based approaches ever use the intermediate reasoning steps provided in the GSM8K dataset, our human annotations, or any models derived from this data. When initializing from the SFT model, we follow Polu and Sutskever (2020) and reuse expert samples from each iteration, so that our training set grows each epoch. We do not do this with few-shot approaches because in that setting, the samples from the early epochs have many trace errors which we do not want the RL model to imitate. Correspondingly, there are several minor implementation differences between the two cases, which we note throughout our detailed descriptions of the policy improvement and distillation procedures.
We consider three versions of the policy improvement procedure (Figure 2). In the Final-answer RL approach, also called Self-taught Reasoner and proposed by Zelikman et al. (2022), we generate full traces per problem and filter by final-answer correctness. For the few-shot version, we select all traces yielding the correct final answer, while for the SFT-based version, we only use one randomly chosen sample per problem. In the ORM-RL approach, we generate full traces per problem, and select the sample with the highest score according to the ORM model. In the PRM-RL approach, we instead treat each step as an individual episode. At each step, we generate candidate steps, select the candidate with the highest PRM score, and continue from the selected step until the model outputs a step with the final answer indicator text, or a maximum of 15 steps. We set across all experiments. For few-shot-based approaches, we retrain the RM after every expert iteration. For SFT-based approaches, we skip this step and use a fixed RM, since somewhat surprisingly, this did not make a significant difference in preliminary experiments.
For distillation, we use the same hyperparameters as SFT. As with SFT, we apply early stopping by validation loss, where our validation set is constructed from expert policy samples on the validation set. For SFT-based approaches, we initialize with the SFT parameters at each distillation step, while for few-shot-based approaches, we initialize with the base model parameters.
7 Data annotation
As discussed in Section 2.4, the PRM is trained on stepwise labels indicating whether the steps so far are correct. To collect this data, we present human annotators with the problem statement, the reference solution from GSM8K, and the generated model solution, and ask them to indicate the first model step with a major mistake, if any exist. Our instructions define a major mistake as “a step where the information expressed is incorrect, or it would no longer be possible to reach the correct solution without undoing that step”. From these annotations, we can label every step with a binary label indicating whether the steps so far are correct: all steps before the first major mistake are labeled ‘correct’, while the remainder are labeled ‘incorrect’.
We applied a small amount of dataset cleaning by removing samples from annotators with low inter-annotator agreement (measured on the 20% of solutions where we used duplicate labelling), as well as those from GSM8K problems flagged by annotators as ambiguous. This removed about 20% of our data, leaving annotations for 1560 model samples across 530 training set problems, corresponding to 9856 step-level binary labels. For the validation set, we used the same procedure, but added duplicate labelling and a manual pass by the paper authors to resolve inter-annotator disagreements. Our validation set contained 162 model samples, with 913 total steps. For evaluation, we used 200 problems with correct final answers per model. This was done for each of the 10 models in Table 1, again with duplicate labelling. We describe full details of our data collection procedure in Appendix B.
Results
Our results are summarized in Table 1 and Fig. 4. The ORM-RL and PRM-RL models achieve a final-answer error rate below 13%, improving on the 16.8% final-answer error for the current state-of-the-art model (Li et al., 2022). This is further reduced to 2.7% when the model is allowed to abstain on only 30% of questions. The corresponding trace errors are 3.4% and 3.8%, which significantly improve on the 14% reported by the best prior work (Wang et al., 2022; Wei et al., 2022). Beyond these quantitative results, we highlight three key takeaways:
The SFT and Few-shot+Final-Answer RL models attain similar final-answer error rates both without an RM (22.3% vs. 23.5%) and with an ORM (14.8% vs. 16.6%). This is notable, as Few-shot+Final-Answer RL only requires demonstrators to provide a final answer, rather than a full reasoning trace. Put another way, Few-shot+Final-Answer RL uses 1-4 tokens of label supervision per question, while SFT uses hundreds. This suggests that in cases where final-answer correctness is sufficient, outcome-based approaches can provide a label-efficient approach with competitive performance.
Despite the fact that ORMs are only trained to predict whether the final answer is correct, we can see in Fig. 4 that ORM predictions tend to agree more with the PRM labels than with the ORM labels themselves (85% vs. 77% averaged over all steps).Appendix C shows similar results when considering just the final step. We suspect this is because it is simpler for the ORM to learn to recognize when steps are correct, than it is to check the answer by internally computing the final answer itself. This is further supported by that fact that, even though trace error is measured only on samples with the correct final answer, RM reranking significantly improves trace error relative to SFT alone (4.4% vs. 11.4%). This suggests that RMs are checking the reasoning steps, and not just the final answers. However, we caution against over-generalizing: the fact that the ORM model approximates the PRM labels may be domain-specific. This may depend on both the relative difficulty for the model to compute the correct answer with and without reasoning traces, and the lack of spurious solutions (Goldman et al., 2017) in math problems, where incorrect reasoning steps are unlikely to lead to the correct final answer.
Fig. 4 shows that despite similar final-answer error rates, there is a significantly higher trace error rate for the outcome-based Few-shot+Final-Answer RL vs. the process-based SFT model (19.8% vs. 11.4%). This discrepancy persists with RM reranking: Few-shot+Final-Answer RL with ORM reranking underperforms SFT with ORM/PRM reranking (12.4% vs. 4.4%/3.5%). However, we find that when we train the few-shot RL model using an ORM (Few-shot+ORM-RL) rather than training directly against final-answer correctness, the trace error drops significantly from 12.4% to 5.5%, closing much of this gap. We believe this results from the previous finding, i.e. that the ORM is basically learning to emulate the PRM allowing the model to learn from emulated process-based feedback and resulting in relatively low trace error rates.
In the rest of this section we provide a more detailed analysis of our results. We group the main results based on how they use the reward model, covering approaches that use no reward model in Section 3.1, approaches that use a reward model only for reranking in Section 3.2, and approaches that use a reward model both during reinforcement learning and for reranking in Section 3.3. We cover selective accuracy and OOD generalization separately in Sections 3.5 and 3.4.
1 No reward model
Comparing the most process-based to most outcome-based approach, SFT and Few-shot+Final-Answer RL have similar final-answer error rates, but SFT has significantly better trace error. Further, when starting from the SFT model, applying Final-Answer RL does decrease final answer error (22.3% to 20.2%), but increases trace error (11.4% to 12.1%, though we note the difference is not statistically significant). These both support the view that outcome-based approaches can find ways to produce correct answers for incorrect reasons. Table 2 provides a qualitative example.
Our focus on different approaches to supervising LMs assumes that, provided sufficient data, finetuning outperforms prompting alone. We first validate this assumption for GSM8K. We find that the 5-shot prompted policy alone (Few-shot) achieves 41.5% final-answer error with majority voting (and 77.7% error when using a single sample). While this is impressive considering that the few-shot policy requires no additional finetuning data, it leaves significant performance on the table compared to both the Few-shot+Final-Answer RL and SFT models, thus validating our initial assumption.
2 Reward model for reranking only
Overall, reward models provide a significant boost to both trace and final answer accuracy. We find that RM reranking significantly improves trace error, reducing it from 11.4% to below 5% for SFT. RM reranking also benefits the Few-shot+Final-Answer RL model, though trace error remains significantly higher than in the SFT case. As found by Cobbe et al. (2021), our results also show that RMs decrease final-answer error, from 22.3% to below 15%. We also experimented with the “step-level voting verifier” from Li et al. (2022). This works similarly to the PRM approach, but computes labels using a heuristic based on intermediate numeric results. This resulted in a slightly worse 15.9% final-answer error.
3 Reinforcement learning with a reward model
We can see from Table 3 that when starting from a few-shot model, RL cuts the final answer error rate by half, regardless of the decoding method. In contrast, when starting with an SFT model, RL has very little effect on top of using an RM for decoding, though does provide a significant improvement when using greedy decoding (41.1% to 31.2%).
In both the few-shot and SFT settings, ORM-RL and PRM-RL outperform Final-Answer RL across all three decoding strategies. On face, this may be surprising given that ORM-RL optimizes an approximation (the ORM full-solution scores) of final-answer correctness. However, our earlier RM analysis (Fig. 4) suggests that the ORM approximates process-based feedback, and checks reasoning steps rather than the final answer directly. Thus, one potential explanation is that Final-Answer RL only checks that solutions reach the correct final answer, whereas PRM-RL and ORM-RL check for solutions which reach the right answer for the right reason.
4 Selective prediction
In many practical applications, it is possible to abstain. For instance, if an ML system was used to explain a problem to a student, or assist in a calculation, it would be preferable for the model to abstain rather than produce an incorrect output. This motivates the selective prediction setting (El-Yaniv et al., 2010; Geifman and El-Yaniv, 2017), where the model is allowed to abstain on of inputs, and selective error rate is measured on the inputs where the model does not abstain. To determine which inputs to abstain on, we set a threshold on the RM score of the selected sample, with the threshold determined by .
Figure 5 shows that by abstaining on % of inputs, we reduce final-answer error rate from 14.1%2.7%, which can be further reduced to 1.5% at %. Further, at , selective prediction with SFT and the ORM or PRM yields a 5 final-answer error reduction, compared to 3 in the Few-shot+Final-Answer RL case. This may be related to the improved trace error for SFT: when trace error is lower, the RM can more reliably use intermediate step correctness to determine its confidence. However, further investigation would be necessary to properly understand this effect.
5 OOD generalization
To evaluate out-of-distribution generalization, we evaluate our models zero-shot on the pre-algebra split of the MATH (Hendrycks et al., 2021) dataset. We limited to problems without Asymptote diagrams, and removed some LaTeX formatting with simple regular expressions, e.g. converting to 3/4. Overall, we see noticeable OOD generalization, with final-answer error rate of 64.6% for our SFT+ORM-RL model with majority voting (74.2% error with no filtering, assuming the model fails all questions with Asymptote diagrams). This is significantly worse than the 29% error on pre-Algebra questions from Lewkowycz et al. (2022), which uses a much larger base LM and trains on more math data, but significantly better than the previous best result of 92.3% error for GPT-3 from Hendrycks et al. (2021). All models have final-answer error in the 60%-70% range, and we do not observe noticeable trends based on the type of supervision. We provide full details and results in Appendix D.
Discussion
Process- and outcome-based feedback have different strengths, and the appropriate choice will often be context-dependent. As a general rule, outcome-based feedback tends to be appropriate when a reliable and complete evaluation metric is available, while process-based feedback is most appropriate otherwise. Here we discuss the considerations that lead us to this view. We start by discussing relative strengths and weaknesses with respect to final-answer and trace error, before moving to motivations for process-based feedback that are yet to be empirically validated.
Our results suggest that when low final-answer error is sufficient, outcome-based approaches provide a label-efficient method for obtaining this, whereas when low trace error is desired, it is helpful to use process-based feedback, or an approximation of it. The relative importance of these metrics is context-dependent. In a context where desired outcomes are easy to evaluate and the evaluation process is robust, final-answer error is appropriate. For example, if final answers can be quickly validated, either programmatically or by quick user inspection, then low final-answer error is a fairly complete performance measure. In contrast, Menick et al. (2022) provide question-answering as an example where low trace error is necessary, since even if the final answer is correct, it is very difficult for users to rely on this answer without also seeing the sources which led to that answer. In other domains such as education, the reasoning steps themselves may also be of direct interest.
1.2 Process-based approaches both require and facilitate human understanding
Relative to outcome-based approaches, process-based approaches tend to require greater human understanding (Krakovna et al., 2020). For example, outcome-based feedback based on power-consumption, chip area, and other metrics can be used to optimize computer chip layouts (Mirhoseini et al., 2021), while a process-based approach would require detailed expert knowledge on designing chip layouts. In order to be competitive in general, process-based approaches require us to improve human understanding, either aided by ML systems such as in Amplification or Debate (Christiano et al., 2018; Irving et al., 2018) or through broader means, such as by working with experts (Rauh et al., 2022), training people (Stiennon et al., 2020), or providing them assistive tools.
Second, process-based approaches may facilitate human understanding because they select for reasoning steps that humans understand. By contrast, outcome-based optimization may find hard-to-understand strategies, and result in less understandable systems, if these strategies are the easiest way to achieve highly-rated outcomes. For example in GSM8K, when starting from SFT, adding Final-Answer RL decreases final-answer error, but increases (though not significantly) trace error.
1.3 Process-based approaches avoid tampering incentives
A common concern within the AI safety literature is from RL agents which tamper (Everitt et al., 2017), i.e., corrupt their feedback mechanisms in order to receive positive feedback. As a hypothetical example, consider an assistive agent which repeatedly interacts with users. If optimized for total user satisfaction, such an agent may influence users towards preferences which are easier to satisfy (e.g. easier to predict, or generally more amenable to ML-generated proposals), in order to increase user-reported satisfaction (Kenton et al., 2021). A similarly long-term concern is with agents which gain power and take complete control of their feedback procedures in order to ensure positive feedback (Cotra, 2022).
In contrast, consider training from process-based feedback, using user evaluations of individual actions, rather than overall satisfaction ratings.Note that the outcome-based vs. process-based spectrum applies in general to supervising any sequence of actions based on the resulting outcomes, or based on each action individually. Because the only actions in GSM8K are reasoning steps, we typically refer directly to reasoning steps in this paper, but the different approaches to supervision apply more generally. While this does not directly prevent actions which influence future user preferences, these future changes would not affect rewards for the corresponding actions, and so would not be optimized for by process-based feedback. We refer to Kumar et al. (2020) and Uesato et al. (2020) for a formal presentation of this argument. Their decoupling algorithms present a particularly pure version of process-based feedback, which prevent the feedback from depending directly on outcomes.
An alternate approach to avoiding tampering is to continually improve the outcome-based metrics, by monitoring for sequences of actions leading to tampering, and penalizing these sequences. However, this approach only scales to tampering we can detect: in other cases, effects can be difficult to observe or measure (e.g., how would we determine if a system influences user preferences?) or understand (e.g., innocuous-looking decisions may still have large effects, particularly in aggregate). Broadly, a risk with the incremental outcome-based approach is addressing the most easily-noticed problems, without addressing root causes or subtler cases.
2 Limitations to generalizability of our results
We generally expect process-based and outcome-based feedback to align more closely for math compared to other domains. For math problems, incorrect traces are typically harmful for reaching correct final answers. This matches our earlier finding that outcome-supervised RMs approximate process-based feedback. In contrast, in other domains, undesirable behaviors may be helpful for highly-rated outcomes, e.g., manipulation may increase reported user satisfaction. As a result, we believe optimizing for outcomes (final-answer correctness) for math problems has a stronger effect on inducing a correct process than it would in other domains.
3 Concepts related to process- and outcome-based feedback
In this work, we focus on different forms of supervision available for training LMs, and discuss these forms of supervision in terms of a distinction between process- and outcome-based approaches. This framing has been used in blog posts (Stuhlmüller and Byun, 2022) and informal discussion, though to our knowledge this is the first empirical paper to discuss it. Here, we discuss similarities and differences to other related distinctions used throughout the literature.
Broadly, supervised approaches tend to be more process-based, and RL approaches tend to be more outcome-based. Indeed, the most process-based approach we consider is purely supervised (SFT), while the most outcome-based approach is pure RL (Few-shot+Final-Answer RL). However, RL approaches can be more or less process-based, such as when comparing PRM-RL to ORM-RL or Final-Answer RL. Conversely, supervised imitation of reasoning traces filtered by outcomes (e.g., traces which led to highly-rated user interactions) blurs together with RL approaches based on supervised learning of high-return trajectories (Silver et al., 2017; Anthony et al., 2017; Abdolmaleki et al., 2018). The process- and outcome-based categorization thus acknowledges that the resulting model depends on the broader data-providing procedure, which often will not be fully described by code.
While the meanings of strong vs. weak supervision can vary depending on context, approaches which supervise intermediate steps are often referred to as strongly supervised (Yang et al., 2018; Perez et al., 2020). Similar to the above, while process-based approaches tend towards strong supervision, and outcome-based approaches towards weak supervision, strongly supervised approaches can be either process- or outcome-based. For instance, a person evaluating intermediate reasoning steps could either be directly evaluating those reasoning steps, or deferring to evaluations of their resulting outcomes.
Verbalized reasoning traces do not necessarily imply process-based approaches. Indeed, whereas all approaches in this work use verbalized reasoning traces, they use both process- and outcome-based feedback. Additionally, the verbalized reasoning traces do not necessarily represent the model’s internal reasoning process, with the bulk of the reasoning happening inside of the large neural network activations, except in strictly modular approaches (Creswell et al., 2022). This caution holds particularly for domains with the potential for deception (Kenton et al., 2021).
Nonetheless, verbalized reasoning still helps enable process-based feedback. In contrast to approaches based on iterative computations performed in activation space (Guez et al., 2019; Dehghani et al., 2018; Graves, 2016; Schrittwieser et al., 2020), humans can directly supervise iterative reasoning steps in natural language.
Related work
Math word problems have been a popular domain for studying reasoning in LMs (Kushman et al., 2014; Ling et al., 2017; Amini et al., 2019; Miao et al., 2020; Hendrycks et al., 2021; Cobbe et al., 2021). Several recent papers have demonstrated that few-shot prompting alone can lead to impressive performance on GSM8K (Chowdhery et al., 2022; Lewkowycz et al., 2022; Wei et al., 2022; Wang et al., 2022). All of these papers include the reasoning traces in the few-shot prompts which encourages the model to generate verbalized reasoning steps (also referred to as self-talk (Shwartz et al., 2020) and chain-of-thought prompting (Wei et al., 2022)). Prompting benefits significantly from training on a large dataset of mathematical content (Lewkowycz et al., 2022), as well as from finetuning for instruction following (Ouyang et al., 2022). Kojima et al. (2022) and Li et al. (2022) demonstrate improvements in final-answer error rate from for GPT-3 (Brown et al., 2020) to for InstructGPT (Ouyang et al., 2022), and further to 23.3% for Codex (Chen et al., 2021).
We focus on finetuning because we are interested in the effects of different feedback procedures, and because it significantly outperforms prompting alone for our base LM. The original GSM8K paper (Cobbe et al., 2021) demonstrated significant benefits of reward models or verifiers, and we use their ORM approach. Li et al. (2022) also study RMs and propose a heuristic-based step-aware RM, which slightly degrades performance on GSM8K, but boosts performance on a wide range of other benchmarks. We find that human evaluations of each step provide an improvement. We also use STaR (Zelikman et al., 2022) (referred to as Few-shot+Final-Answer RL throughout this paper), and show its GSM8K final-answer error can be reduced from their reported 89% to 23.5% through the use of a better base model (Hoffmann et al., 2022) and further reduced to 13.8% by using an RM-based RL instead of their final answer RL procedure. In contrast to the above prior work, we not only show improved performance, but also provide a comprehensive comparison across different types of feedback, with a focus on trace error rate in addition to final-answer error rate.
Moving beyond math problems, a large body of work studies multistep reasoning for LMs. While a full review is beyond the scope of this work, we discuss a few representative categories of approaches. Prior work has suggested improvements to the base model (Lewkowycz et al., 2022; Ouyang et al., 2022), as well as prompt-based approaches (Perez et al., 2020; Shwartz et al., 2020; Wei et al., 2022; Kojima et al., 2022; Dohan et al., 2022). Our focus on supervision techniques for finetuning is complementary to such improvements.
Other work has focused exclusively either on outcome-based or on process-based approaches. For example, on the process-based side, Wu et al. (2021) summarize full-length books by supervising individual summaries, which are recursively composed. This provides an example where the outcome-based approach (directly using human approval of full-book summaries) would be prohibitively expensive, due to the cost and sparsity of such feedback, but process-based supervision of each step individually is possible. Creswell et al. (2022) and Nye et al. (2021) use process-based SFT in synthetic settings where reasoning traces can be synthesized programatically, while we instead focus on a natural language setting where reasoning traces and feedback must be generated by humans. On the outcome-based side, Zelikman et al. (2022) apply an outcome-based approach to multi-step reasoning questions. In contrast to all of this work, however, we directly compare outcome-based and process-based techniques, and include a detailed analysis on trace error rates.
To our knowledge, the prior work which most directly compares process- and outcome-based feedback for multistep LM reasoning is WebGPT (Nakano et al., 2021). In WebGPT, intermediate steps are web browser interactions, which they supervise either via SFT, or via outcome-based RL on resulting answers. Similar to us, they observe significant improvements over SFT from outcome-based RM reranking, but adding full-solution RL on top has minimal effect. However, we more comprehensively explore both process- and outcome-based supervision, additionally evaluating process-supervised RMs, the PRM-RL approach, and purely outcome-supervised RL policies (without SFT). This allows us to draw broader conclusions regarding the effects of supervision on final-answer vs. trace error compared to prior work.
Most prior work on head-to-head comparisons of process- and outcome-based approaches has been on algorithmic tasks such as sorting lists of numbers, which avoid working with human data. Early outcome-based approaches such as Neural Turing Machines (Graves et al., 2014) are trained end-to-end to predict the final answer. In contrast, Neural Programmer-Interpreters (Reed and De Freitas, 2015; Li et al., 2016; Cai et al., 2017) are trained to imitate each step in an execution trace, which can then be chained together. This requires stronger supervision but improves generalization. Iterated Amplification (Christiano et al., 2018) uses a bootstrapping procedure to train models to approximate a potentially exponentially sized tree of reasoning steps, assuming an oracle that cannot directly answer hard problems, but can decompose them into easier ones. Our work extends these results to natural language, where neither programmatic execution traces nor decomposition procedures are available, and these must be learned from human feedback.
In this work we chose to work with the GSM8K dataset, because it provides natural language reasoning traces that allow for a detailed comparison between process- and outcome-based approaches without requiring us to collect the traces ourselves. Alternative datasets which also contain full reasoning traces include EntailmentBank (Dalvi et al., 2021), StrategyQA (Geva et al., 2021), ProofWriter (Tafjord et al., 2020), and CLUTTR (Gontier et al., 2020). However, compared to GSM8K, these datasets either contain templated problems (ProofWriter and CLUTTR) or are significantly smaller in size (EntailmentBank and StrategyQA). Furthermore, working with multiple datasets would be expensive because of the need to train human annotators for each task, and to collect a significant number of human feedback annotations for each dataset.
Conclusion
In this work, we run the first comprehensive comparison between process- and outcome-based supervision on a natural language task. We find that both types of supervision lead to similar final-answer error rates, with our best models improving the state-of-the-art final-answer error on GSM8K from 16.8% to 13.8% when using outcome-based supervision and to 12.9% when using process-based supervision. In contrast, we find that obtaining low trace error requires either process-based supervision, or a reward model that emulates it. A purely process-based approach of SFT with PRM reranking reduces the state-of-the-art trace error rate from 14.0% to 3.4%, while its outcome-based analogue achieves 12.7% trace error. However, somewhat surprisingly, we find that reward models trained with outcome-based labels result in predictions that agree more closely with the process-based labels than they do with the outcome-based labels themselves. By using this reward model during RL training, we close most of this gap, reducing trace error from 12.7% to 5.5%. While some of these conclusions may be specific to our setting of math word problems, we hope that future work explores the extent to which they generalize to other domains.
References
Appendix A Example GSM8K problems and solutions
We include several examples to provide a qualitative sense of the task and learned model behavior. LABEL:tab:random_samples contains 10 randomly sampled problems, and the output of the SFT+ORM-RL model with ORM reranking. LABEL:tab:trace_errors contains 5 trace errors, where the final answer is correct, but at least one of the steps has an error.
Appendix B Data annotation details
The full details of our study design, including compensation rates, were reviewed and approved by DeepMind’s independent ethical review committee. All participants provided informed consent prior to completing tasks and were reimbursed for their time. It is our policy that researchers must pay workers/participants at least the living wage for their location.
For constructing the PRM training dataset, we used samples from the SFT model. Due to a limited annotation budget, we only annotate problems where the SFT majority voting prediction is incorrect, since this focuses training on difficult problems, and still includes a mix of correct and incorrect samples due to annotating 3 model solutions per problem.
Since evaluating the accuracy of solutions to mathematical problems is a special skill, we ran a preliminary qualification study before using annotations from participants for training or evaluation. For the qualification annotation tasks, we selected model solutions where three authors unanimously agreed on the first major mistake, and required participants to annotate at least 3 out of 4 such solutions correctly. In total, 21 / 91 candidates were included in our annotator pool.
Additionally, we used duplicate annotations for 20% of the training problems. We removed ratings from annotators who had an inter-annotator agreement rate below 75% on doubly-rated problems. On manual inspection, we found that these annotators had typically made errors in this disagreeing cases. This removed data from 4 / 21 annotators, amounting to 21% of our originally labelled training data, leaving us with 530 annotated problems (compared to 675 originally).
After quality assurance steps, using the same set of duplicated problems, we measure inter-rater agreement rate of 92% and Cohen’s of for the task of predicting the first incorrect step. Note that this estimate will be slightly biased upwards as we filtered raters based on this same set, but as the inter-rater agreement rate is fairly bimodal across raters, this effect should be relatively small.
For the evaluation, inter-rater agreement is 87%, with Cohen’s of on the binary task of labelling the full trace as correct. Inter-rater agreement is significantly lower than on the training set. We attribute this to the observation that for solutions with an incorrect final answer, there is often a fairly clear-cut mistaken step, but for cases with the correct final answer, whether or not an intermediate step is incorrect can be more subtle.
Appendix C Additional RM analysis
Figure 6 shows the agreement matrix between RMs and RM labels on the last step. Similarly to Fig. 4 (which averages over all steps, rather than just the last step), we see that the ORM has higher agreement with the PRM labels, despite being trained to predict the ORM labels. Note that on the last step, the ORM and PRM labels will exactly coincide, except for the case where the final answer is correct, but the trace is incorrect (by construction, it is impossible for the trace to be correct with an incorrect final answer). Thus, this result indicates that the ORM tends more towards predicting whether the full trace is correct, and not just whether the final answer is correct.
Appendix D OOD evaluation details
To measure out-of-distribution generalization, we evaluate our models zero-shot on the pre-algebra split of the MATH (Hendrycks et al., 2021) dataset. We include results for additional models in Table 7, showing noticeable OOD generalization. An example question and answer from the dataset can be see in Table 6. We limited to problems without diagrams, leaving 633 of the original 871 problems, and removed some LaTeX formatting as discussed below.
/\3 \verbdef\secondcap\2 (2nd capture group)
We used a simple regular expression (regex) based transformation to convert Latex math expressions to plain texts, since the GSM8K data is not in Latex format and we were more interested in OOD generalization across problems rather than formats. See Table 8 for an example before and after conversion. Our goal was not to cover all Latex commands and formatting edge cases in the MATH dataset, but instead we wanted to address the more commonly used commands for which we could write regular expression transformations. See Table 9 for the full set of transformations we applied.
In the MATH dataset, the final answer is always in a box, by formatting the answer with the “\boxed” command, e.g. . To obtain the final answer from the solution text, we find the boxed expression and extract its argument.
Appendix E Negative and preliminary results
Here, we include various observations we made while running experiments. These are much less carefully checked than the results reported in the main paper, and some may be specific to details or our training setup. However, we include them as we believe they may nonetheless be helpful for future researchers:
Training the ORM required roughly 20 times more steps than SFT training. We believe this is likely due to a combination of sparser token supervision (one token per step, rather than hundreds of tokens per solution), requiring dropout regularization, and being able to train on multiple generated samples per problem.
We experiment with step-level reranking, by applying an argmax over options at each step, rather than solution-level reranking. However, we found this increased final-answer error by about 1% with the PRM, and 3% with the ORM. From a manual inspection by the first author of problems solved by solution-level reranking, but not step-level reranking, this seemed to largely be due to insufficient entropy in the policy. In these failure cases, there would often be a step where a valid next step existed, but none of the continuations found it (and often, all continuations would make the exact same mistake). In contrast, with solution-level reranking, slight variations in the earlier generated steps would yield more variation across the corresponding completions. This could potentially be addressed by, for example, conditioning the model on steps rejected by the RM, but for simplicity we used solution-level reranking throughout this paper.
In preliminary experiments, retraining the ORM between expert iterations during SFT+ORM-RL training did not help. In contrast, for the few-shot-based expert iteration implementations, we do retrain the ORM every iteration, in part to avoid imitating trace errors in early epochs.
During training set collection, we initially asked annotators to provide a corrected version of the mistaken step. However, it was difficult to communicate this task in sufficient precision: for example, corrected steps would include calculation outputs before the calculations (which are difficult for an autoregressive LM to handle), or would try to stay too close to the reference solution even when the earlier model steps were part of a different (but still valid) approach. Additionally, unlike for the labelling task, we couldn’t easily compute inter-rater or researcher-rater agreement rates.