Scaling Relationship on Learning Mathematical Reasoning with Large Language Models
Zheng Yuan, Hongyi Yuan, Chengpeng Li, Guanting Dong, Keming Lu, Chuanqi Tan, Chang Zhou, Jingren Zhou
Introduction
Large language models (LLMs) (Anil et al., 2023; Touvron et al., 2023b; OpenAI, 2023) have shown considerable abilities in various math reasoning tasks (Saxton et al., 2019; Cobbe et al., 2021; Lightman et al., 2023). It is of interest to understand, predict, and improve an LLM’s math reasoning ability based on different pre-trained LLMs and supervised datasets. With this knowledge, we can better decide the effort we put into improving the LLM or augmenting the dataset. Many recent works are focusing on using different prompts (Wei et al., 2022b; Yao et al., 2023) or ensembling / reranking multiple times of inferences (Cobbe et al., 2021; Uesato et al., 2022; Wang et al., 2023; Lightman et al., 2023) to improve models’ reasoning performances. While in-context learning (ICL) and performing multiple inferences can improve performance, it is computationally expensive and not suitable for online deployment scenarios. Therefore, we focus on the performance of the supervised LLMs with inference only once which is a setting closer to online deployment.
To this end, we empirically investigate the scaling relationship of factors that influence the math reasoning abilities of a supervised LLM, including pre-training losses, the amount of supervised data, and the amount of augmented data. Firstly, we analyze the supervised fine-tuning (SFT) and ICL performance of LLMs. We observe that the pre-training loss is approximately negatively linear correlated to the SFT and ICL accuracy in a given interval which is a better performance indicator than pre-trained model sizes or pre-trained token counts. Secondly, we analyze the relationship between SFT and different amounts of supervised data. We observe that the model performance has a log-linear relation versus the supervised data amount while the increase diminishes with the better pre-trained model. Thirdly, we want to leverage the model itself to generate more supervised data to reinforce its reasoning ability and analyze the scaling relationship of the augmented data amount. We apply rejection sampling on SFT models to sample and select correct reasoning paths as augmented dataset (Uesato et al., 2022; Zhu et al., 2023). We use these augmented datasets to fine-tune base LLMs which would achieve better performances compared to SFT and we denote it as rejection sampling fine-tuning (RFT). We find the key factor influencing RFT performance is the distinct reasoning path amount which can be increased by sampling more times or combing samples from multiple models. We apply RFT on several pre-trained LLMs and show larger improvement on less performant models. We discuss the reason RFT works is it provides multiple reasoning paths which makes LLMs have better reasoning generalization. We also discuss that RFT is much cheaper than pre-training in computational resources while training an LLM with lower pre-training loss is the fundamental solution.
The key findings of this paper are shown in Figure 1 and are summarized here:
When the pre-training loss gets smaller (i.e. the pre-trained model gets better), the model reasoning performances of SFT and ICL increase linearly within a range. The SFT performance improves slower than ICL.
SFT improves in a log-linear manner with the increase of supervised data amount. The benefits of increasing data amount diminish as the pre-trained model gets better.
The model performance for RFT improves as the distinct reasoning path amount increases. The RFT performance improves slower than SFT.
The combination of rejection sampling samples from multiple models further enhances the RFT performance, resulting in an accuracy of 49.3 for LLaMA-7B (+13.4 compared to SFT), 50.3 for LLaMA2-7B (+8.7 compared to SFT), 52.1 for LLaMA-13B (+9.1 compared to SFT), and 55.4 for LLaMA2-13B (+5.4 compared to SFT).
Related Works
Recent research on LLMs has discovered the emergent ability to solve reasoning tasks beyond a certain model scale (Wei et al., 2022a). Such reasoning abilities in LLMs can be elicited by fine-tuning, few-shot prompting, or zero-shot prompting (Cobbe et al., 2021; Wei et al., 2021; Nye et al., 2021; Wei et al., 2022b; Kojima et al., 2022). A large amount of research focuses on the reasoning tasks of math word problems (MWP), and methods are evaluated on the benchmarks spanning different levels of MWPs (Koncel-Kedziorski et al. (2016); Patel et al. (2021); Lan et al. (2021); Cobbe et al. (2021); Jie et al. (2022); Yuan et al. (2023a); Fu et al. (2023a), inter alia). The core idea of improving the mathematical reasoning ability of LLMs is to aggregate various sampled reasoning paths during either fine-tuning or inference. Cobbe et al. (2021) trained and devised a reasoning path verifier to select the correct results during inference. Wang et al. (2023) proposed to sample various reasoning paths during inference and then derive the final result by majority voting on the answers or through verifiers (Li et al., 2023). Several works applied the idea of rejection sampling along with other techniques to filter the diverse sampled reasoning paths for fine-tuning data augmentation (Huang et al., 2022; Zelikman et al., 2022; Ni et al., 2023; Zhu et al., 2023). Rejection sampling is a simple-yet-effective fine-tuning augmentation technique and is also used for LLM alignment with human preference (Bai et al., 2022; Yuan et al., 2023b; Dong et al., 2023; Touvron et al., 2023b; Song et al., 2023). Uesato et al. (2022) explored to use of reinforcement learning methods for improving the mathematical reasoning abilities of LLMs and they further discussed the difference between outcome-based and process-based reward modeling. Followed by Lightman et al. (2023), they collected large-scale process-based supervision signals through human annotation and verified that LLMs can benefit more from process-based reward modeling with human-annotated supervision than outcome-based reward modeling. There is also prior research that distilled the emergent reasoning ability of LLMs to small language models (Fu et al., 2023b; Shridhar et al., 2023). Compared to previous works (Zelikman et al., 2022; Uesato et al., 2022; Zhu et al., 2023; Ni et al., 2023), we are using a simpler way of generating augmented samples without any trained process-level reward models and we are focusing on researching the scaling relationship between LLMs and math reasoning ability.
Scaling Laws of Large Language Models
It is important to understand and predict the performance gain as the language model scales up. Kaplan et al. (2020) first investigated and derived a predictable relationship on how the number of model parameters and data sizes contribute to the loss over many orders of magnitudes. Hoffmann et al. (2022) refined the scaling laws in (Kaplan et al., 2020) and found the scaling laws for computation-optimal training. Muennighoff et al. (2023) explored and extended the scaling laws under a data-constrained scenario. Besides investigating the scaling performance for pre-training, Gao et al. (2022) discussed the scaling laws for overparameterized reward models for alignment with human preference, and Hernandez et al. (2021) developed scaling laws for transferring performance from pre-trained models to downstream tasks. Henighan et al. (2020); Caballero et al. (2022) investigated scaling laws of math problems. In this paper, we are investigating the scaling relationships of large language models on learning math word problems with pre-training losses, supervised data amount, and augmented data amount.
The factors of math reasoning ability in Supervised LLM
The target of this paper is to try to understand the performances of supervised LLMs in math reasoning. We expect a pre-trained LLM to learn reasoning ability from a supervised reasoning dataset . The dataset is defined by , where is a question, is a chain-of-thought reasoning path, and is a numerical answer. We perform supervised fine-tuning on dataset to obtain an SFT model . We use to generate reasoning paths and answers in the test set by greedy decoding and report the accuracy (i.e. or maj1@1) as our metric here.
Previous works state that the larger LLM shows better reasoning ability across the same series of models (Brown et al., 2020; Chowdhery et al., 2022; Touvron et al., 2023a; b), and we find LLaMA outperforms GPT-3 which shows the model parameter counts should not be the only indicator of reasoning ability. While LLMs have different architectures, model parameters, and pre-training token numbers, we find the pre-training loss is a stable performance indicator of the math reasoning ability and we use it to represent the model instead of using their model parameters and pre-training token numbers.
We analyze the SFT and ICL (8-shot) performance of GPT-3 (Brown et al., 2020), LLaMA (Touvron et al., 2023a), LLaMA2 (Touvron et al., 2023b), and GPT-4 (OpenAI, 2023). The pre-training losses of these models are observed in their paper, we should notice that pre-training losses correspond to different pre-training datasets and different tokenizers which means they could not be compared strictly (and we cannot use it to do any sort of regression directly) while the tendency among these losses is still enlightening. We use the results of GPT-3 fine-tuning from (Cobbe et al., 2021) and we fine-tune LLaMA and LLaMA2 on the GSM8K training set (detailed in Appendix A.1). For in-context learning, we use the results from LLaMA (Touvron et al., 2023a) and LLaMA2 (Touvron et al., 2023b) paper.
The pre-training losses are approximately negatively linear correlated to the SFT and ICL accuracy during the given pre-training loss interval.
SFT outperforms ICL consistently, while the improvements diminish when the pre-training loss is lower.
The linear relation of SFT and ICL accuracy may only work in the given interval. The reasons are (1) the slope of ICL is steeper than SFT, while the SFT performance should be greater than ICL performance; (2) the accuracy can not bigger than 1 or smaller than 0. It should be using instead of as the dependent variable theoretically while we find an apparent linear relationship among pre-training loss and and use as the dependent variable. LLaMA-2 7B(13B) can be viewed as an approximation of continue-training of LLaMA 7B(13B). As it trains longer, its ICL and SFT performance both improve without changing the parameter count. From the observations, one effective way to improve reasoning ability is to train a better base model with lower pre-training loss (Pre-training is all you need!). The models with lower pre-training loss improve less from the fine-tuning which may be due to the models having already obtained more reasoning abilities during pre-training and the supervised data can provide less signal to supervise them.
2 Model Accuracy vs. Supervised Data Count
Supervised fine-tuning does improve LLMs’ reasoning ability, we want to know how the supervised data amount influences the model’s improvement. We fine-tune LLaMA and LLaMA2 with amount of the training set from GSM8K (detailed in Appendix A.2). We want to use this experiment to extrapolate the model performances if we have more supervised data. In Figure 3, we plot the results of training with different amounts of supervised data. From this figure, we can observe that:
The model performance has a log-linear relation versus data amount. When the data amount doubles, the performance increases by a unit.
Better model needs more amount of data to outperform its ICL performance.
Better model benefits less when supervised data amount doubles.
The log-linear relation is stable during amount of the training data. From the observation, it is straightforward to enlarge the training dataset to improve the performance, especially for worse models. For better models, it benefits less which echoes that better models have learned more reasoning ability during pre-training.
3 Model Accuracy vs. Augmented Data Count
Increasing the amount of math reasoning labeled data is difficult, especially proposing a new question. It is easy for a well-educated student to solve hundreds of math word problems per day, but it is very hard to come up with diverse and educational math problems. So our direction changes to augment new data using existing resources. We have tried augmenting new queries (detailed in Appendix D.1) and augmenting revisions (detailed in Appendix D.2). These approaches have none to marginal improvements compared to SFT. We find a simplified version of rejection sampling (Zhu et al., 2023) is a naive and effective way to augment new reasoning paths and can improve the model performance. And we find the key factor influences fine-tuning on rejection sampling (RFT) augmented data is distinct reasoning path amount. Combining rejection sampling samples from multiple models, we can further fine-tune a LLaMA-7B model to an accuracy of 49.3 (compared with SFT 35.9) and a LLaMA-13B model to an accuracy of 52.1 (compared with SFT 43.0).
The SFT model obtains the ability to perform zero-shot chain-of-thought reasoning, and we use to generate more correct reasoning paths to supply the training dataset. For each , we generate candidate reasoning paths and answers with a temperature of 0.7 following (Cobbe et al., 2021). We first filter out reasoning paths with wrong answers or wrong calculations based on Python evaluation. Each reasoning path contains a list of equations , and we select one reasoning path for each distinct equation list as the augmented data and remove other reasoning paths with the same list of equations to deduplicate similar reasoning paths. Different order of elements (e.g. and ) or different order of equations (e.g. and ) are considered different. It is helpful for models to know these orders can be exchanged and is hard for models to learn this with only one reasoning path each problem. We define as the augmented dataset. We fine-tune on pre-trained LLM to as RFT, and we detail how we apply RFT in Appendix A.3. We list the results of RFT with sampling candidate reasoning paths on LLaMA and LLaMA-2 in Table 1. For ICL, SFT, and RFT, we list the maj1@1 (accuracy) and maj1@100 (sample 100 times and calculate accuracy based on majority voting) as metrics.
In the case of 7B and 13B models, RFT yields an approximate increase of 5 to 6 points in maj1@1 and about 4 points increase in maj1@100. For 33B models, RFT does not improve performance compared to SFT. The main reason comes from the augmented samples from rejection sampling. We can find that better models generate more correct reasoning paths per question. For LLaMA-33B-SFT, it can generate an average of 88.7 correct paths per question. However, it overfits the training set and has difficulty generating more diverse paths on the training set questions. Rejection sampling with 33B is very time-consuming and we do not conduct a temperate grid search, we have tried using a larger temperate 1.0 for decoding LLaMA-33B-SFT models, it generates 82.4 correct paths and 4.77 distinct paths per question which is more diverse than using temperate 0.7 but still less diverse than 7B and 13B models. We admit there should be a temperate (or generation config) that can produce more distinct paths and generate good results for RFT in 33B and even larger models while it does need more computation resources for inference compared to sampling using 7B and 13B models. We will show we can use 7B and 13B models only for rejection sampling to improve the 33B model.
Model Accuracy vs Rejection Sampling Data Count
To understand the performance of RFT, we vary among and apply RFT. We also have another setting of while not removing any reasoning paths denoted as no dedup. We list the RFT results with different on Figure 4. Comparing using RFT with and no dedup, the performance is similar and shows that it is better to estimate RFT performance based on distinct reasoning path amount instead of RFT augmented sample counts. Furthermore, using deduplication has better performances for 3 of 4 models and needs much less training time.
When using , RFT outperforms SFT by 2 points stably. For most data points, using larger leads to better performances. However, the merits of RFT are decreasing when doubling . We calculate different paths per question for different in Table 2. We can see that the amount of different reasoning paths is not growing quickly along growing. In Figure 3, we know doubling training samples can have a linear performance improvement. Doubling reasoning paths should improve less than doubling training samples since obtaining different reasoning paths does not obtain any new questions. Therefore, doubling leads to diminished performance improvements.
Combining rejection sampling samples from multiple models
The experiment results above demonstrate performance boosts in mathematical reasoning, benefitting from rejection sampling. Through case studies in 4.1, we show that rejection sampling can augment training data with reasoning paths of diverse calculation processes. However, the reasoning paths sampled from one single SFT model can be logically non-diverse. Therefore, we expect to further improve the mathematical reasoning performance by leveraging rejection sampled reasoning paths aggregated from different models. We denote two final datasets as and , which are aggregated from rejection sampling different models and , where U means models under a certain size, 7B/13B/33B means LLaMA-7B/13B/33B and 7B2/13B2 means LLaMA2-7B/13B. means an aggregation process in which all the reasoning paths from different sets are first combined and then Algorithm 1 is applied to deduplicate the reasoning paths with the same calculation process regarding the equation forms and orders.
We can see, through the results visualized in Figure 5, that using the aggregated dataset and can lead to uniformly better performance than fine-tuning with datasets from a single model across different model sizes. RFT on these two augmented datasets and decreases the performance gaps among the same size models in SFT and RFT which mean the combined augmented datasets provide enough reasoning supervision to fulfill the pre-training gap. We can assume with sufficient supervised data amounts, the performance indicator should be the model size but not the pre-training losses.
We have stated that it is expensive to apply RFT on 33B models and it needs a temperate grid search to achieve an improvement compared to SFT. However fine-tuning on has similar rejection sampling computational cost compared with sampling 100 times on 33B and achieve better performance.
Another phenomenon is including in aggregation barely influences the performance. To give a more comprehensive analysis of the results, we calculate the average reasoning path number per question in Table 2 and depict a Venn diagram to visualize the source of different reasoning paths shown in Figure 6. In Table 2, the average reasoning path numbers of and surpass those of a single model by large amounts, while only have slightly more reasoning paths than by 0.81. In the meanwhile, as shown in Figure 6, the models under and including the size of 13B can contribute unique reasoning paths of similar proportion in around 15%. However, only 6.5% of the reasoning paths can be exclusively acquired from LLaMA-33B-SFT model. This shows that the SFT model of 33B can provide limited reasoning diversity when sampling the training questions. This finding is consistent with the results above in Table 1, indicating the 33B model (and possibly 65B and 70B models) can well memorize the human-annotated reasoning paths.
For 65B models, we find using does not improve the performance compared to SFT. The reason can be better models benefit less from the supervised sample amounts while it has learnt more reasoning ability during pre-training.
Overall, we can come to the conclusion that (1) RFT improves the mathematical reasoning performance of (worse) LLMs through diverse reasoning paths from rejection sampling of the SFT models, and aggregating more diverse reasoning paths can improve the performance further. (2) Different SFT models can contribute reasoning paths with different calculation processes from rejection sampling, leading to more diverse training data for RFT, and LLMs of larger parameter sizes may degrade in generating diversified reasoning paths as a result of overfitting the training questions. There may be a generation config or training config for large enough LMs not to overfit on the training dataset while it is not trivial to find them.
Comparing to other baselines
We compare our RFT results of training on to several baselines and the results are detailed in Table 3. Although LLaMA and LLaMA2 are top-tier open-sourced LLMs https://huggingface.co/spaces/HuggingFaceH4/open_llm_leaderboard, their mathematical reasoning performances still lag behind the current proprietary LLMs which are of larger parameter scales, such as GPT-4 and PaLM2. Compared to results on open-resourced models, our results on LLaMA present better performance than two recent state-of-the-art reasoning augmentation methods. Our RFT method is simpler compared to CoRE, since RFT does not require training verifier models and decoding with Monte Carlo Tree Search (MCTS). Compared to other open-sourced aligned language models, we can find that 7B models struggle at a level of 35 scores which are very similar to SFT performances of LLaMA-7B. We guess they use GSM8K during their pre-training phase following (OpenAI, 2023) or human alignment fine-tuning phase following (Qingyi et al., 2023). Using our augmented dataset to replace the original GSM8K can significantly boost their 7B models’ performances.
Discussion
In the aforementioned analysis of RFT training data, we observe that rejection sampling can augment the training question with diverse reasoning calculation paths. In this section, we investigate whether RFT models can learn to generate different reasoning paths to reach the correct answers. We fine-tune LLaMA and LLaMA2 of 7B and 13B on . During inference, we sample 100 different reasoning paths from each trained model for each test set question with a temperature of 0.7. For each question, we compute the number of different calculation processes presented in 100 sampled reasoning paths that lead to the correct answer and draw histograms with respect to test set questions. SFT and RFT models on self-sampled datasets (RFT k=100) are included for comparison.
As shown in Figure 7, the models trained by RFT on exhibit more question counts than the models trained by RFT k=100 and SFT on the larger numbers of unique calculation processes. There are more question counts for SFT models where all the sampled reasoning paths only correspond to one single calculation process and SFT models can barely generate more than 8 different calculation processes for a question. This analysis demonstrates that diverse reasoning calculation paths in training data can equip the LLMs with finding diverse reasoning logic for solving math problems.
2 Towards Excelsior Mathematical Reasoning
From our findings, there are two main factors that can improve mathematical reasoning abilities given a preset amount of human-annotated samples, including: (1) Pre-training the LLMs to lower losses; (2) Augmenting fine-tuning with rejection sampling. Through extensive experiments, we empirically verify the scaling relationships between the mathematical reasoning performance of LLM with both factors respectively. Out of the consideration of sustainable NLP, in this section, we investigate the possible computational resources required to extrapolate the mathematical performance of LLMs by both factors and discuss how to improve the performance more efficiently.
We estimate the pre-training, SFT, RFT inference, and RFT FLOPs following Kaplan et al. (2020) and GPU times in Table 4 which is detailed in Appendix E. We can find that the cost times of SFT () and RFT () are negligible compared to pre-training. One can always use SFT and RFT to improve models’ performance. However, it could be hard to use RFT to further boost performance. Since we need much more sampling counts (at an exponential level) to increase distinct reasoning paths and there exists an upper bound of distinct reasoning path amount for a given math reasoning question.
We assume that performance follows RFTSFTICL, from the findings in this paper we know the improvement speed follows RFTSFTICL. And if we have an omnipotent language model which has a pre-training loss that is the same as the corpus randomness, it could have RFT = SFT = ICL = 100. Thus when you pre-train a better language model (i.e. smaller pre-training loss), your model’s performance still follows RFTSFTICL but their performance gaps are diminishing. Since you can obtain an RFT model without too much effort (compared to pre-training), then the most important thing we should do is to decrease the model’s pre-training loss. From LLaMA-7B to LLaMA2-7B, it needs to add FLOPs to obtain a 2.1 improvement in the RFT-U33B setting with a 0.05 pre-training loss decrease. From LLaMA-7B to LLaMA-13B, it adds FLOPs to obtain a 2.3 improvement in the RFT-U33B setting with a 0.07 pre-training loss decrease. While minimizing pre-training loss is expensive compared to SFT and RFT, we believe other abilities may follow a similar pattern and better pre-training can benefit all other tasks.
Conclusions
In this paper, we are investigating the scaling relationship in supervising math reasoning abilities with large language models. We find the relationship between math performance and pre-training losses, supervised data amount, and distinct reasoning paths. We find that better language models benefit less with SFT and RFT, and the most important thing is to pre-train a better language model towards excellent math reasoning abilities.
Acknowledgement
We would like to express our sincere appreciation to Tianhang Zhu, Runji Lin, Kai Dang, Keming Lu, Wei Wang, and Junyang Lin for their valuable insights and contributions to this paper.
Limitations
In this paper, we miss the following parts which are very important for building math reasoning abilities for LLMs and should be discussed in the revised version of this paper or future works.
Pre-training on the math-related corpus. This is obviously useful shown in Lewkowycz et al. (2022). While the pre-training loss obtained here cannot align with general domain pre-trained models’ losses.
We do not regress any scaling laws in this paper since many numbers are estimated and pre-training losses, ICL prompts and SFT settings of various models may not be aligned.
References
Appendix A Detailed experiment setting
We fine-tune GSM8K with 3 epochs and a batch size of 128 on NVIDIA A100 GPUs. We use 8 GPUs for 7B and 13B models, 16 GPUs for 33B models, and 32 GPUs for 65B and 70B models during fine-tuning. We use a peak learning rate of 2e-5 with a learning rate warmup. We evaluate the results on the final epoch. We use greedy decode to calculate maj1@1 and decode with temperature 0.7 to calculate maj1@100.
A.2 SFT on downsampled GSM8K
We random downsample GSM8K dataset for fine-tuning. We find that using 3 epochs for little data will result in very poor results which are listed in Table 5. We search training epoch among and evaluate the latest epoch. We report better test results among these two different epoch settings.
A.3 Rejection Sampling Fine-tuning on GSM8K
We use an SFT model to sample on training dataset for times with a temperature of 0.7. We extract the equation list in generated reasoning paths by finding first, removing all white spaces, and joining the equation string list by a special symbol to a string (called it get_equation in our algorithm) for deduplication. We select the reasoning paths by this algorithm:
We are trying to find the most dissimilar reasoning paths based on Levenstein distances. The idea comes from we want diverse reasoning paths for better generalization.
Appendix B Detailed Results of SFT and RFT
We list detailed results of SFT and RFT in Table 5 and 6.
Appendix C Case Study of RFT
In this section, we present the cases of the training samples from rejection sampling. The case studies would shed light on how RFT potentially improves the mathematical reasoning performance of LLMs. The cases are shown in Table 7. As aforementioned, RFT considers the reasoning paths with different calculation processes regarding equation forms or orders, leading to the correct answers. In the cases from Table 7, all the reasoning paths from RFT result in the correct answer of 10, while the calculation processes of reasoning are diverse. Path 1 and 2, as well as Path 4 and 5, are different in the equation forms as highlighted in red. Path 1 and 2 present a two-step calculation reasoning process while Path 4 and 5 alter to a one-step calculation reasoning process. The case demonstrates that rejection sampling can potentially provide more supervision signals that improve mathematical reasoning performance. The filtered reasoning paths sampled from LLMs themselves are of similar quality to the reasoning demonstrations from human annotations.
Appendix D Preliminary Experiments
Through our preliminary experiments and case studies, the errors made by the fine-tuned LLMs are partly attributed to the incorrect reasoning chains where LLMs mistakenly understand the context information or fail to consider all the information in the queries. Although such incorrect reasoning chains lead to wrong answers to the original queries, the reasoning chains themselves represent reasonable logic. For example, for the query Josh decides to try flipping a house. He buys a house for 50,000 in repairs. This increased the value of the house by 150%. How much profit did he make?, a fine-tuned LLaMA model predicts The value of the house increased by 80,000*.15=92,000. So he made a profit of 92,000-80,000-50,000=$42,000 where the model erroneously interprets 150% as 15%, but the reasoning chain is reasonable if we ignore the error.
Therefore, such wrong predictions made by the LLMs may be correct under other queries (if we change 150% to 15% in the above example). We conduct experiments to generate queries for the predicted reasoning chains. This is a similar idea to the hindsight experience replay (Andrychowicz et al., 2017) in reinforcement learning where the method is designed to deal with the sparse reward problems by changing the original objectives for the failed samples to form samples with positive rewards. Such an idea was recently adopted by HIR (Zhang et al., 2023) to better align LLMs with instructions.
Concretely, we reformat GSM8K reversely by predicting the query given the corresponding ground-true reasoning result and then we fine-tune a LLaMA model on the reversed task. We use this model to generate queries on the predicted reasoning chains by a normally fine-tuned LLaMA model on the training set of GSM8K, formalizing a training sample for augmentation. We experiment on the LLaMA 7B model and fine-tune models on the data mixing original and generated samples or solely on generated samples.
The results are shown in the left subfigure in Figure 8. We can see that fine-tuning with self query augmentation data leads to the worst results, and the performance of mixing the original data with self query augmented data still falls short of that of the original data. The fine-tuned performance for mathematical reasoning does not benefit from the naive idea of self query augmentation. Through several case studies of generated data, we find that there are two major defects in the generated data. The first one is some reasoning chains themselves are not logically reasonable, for example, there may be some calculation errors in the reasoning chains. The second one is that the generated query may not be suitable for a reasoning chain. The query generation model may still erroneously interpret the information in the reasoning chains. Both defects attribute to a mediocre augmented data quality, hence can be possible reasons for the failure of this data augmentation procedure.
D.2 Self Revising Augmentation
We also explore improving the mathematical reasoning abilities of LLMs through revising augmentation. To equip LLaMA with revising abilities, we generate a revising dataset by first sampling reasoning paths from a fine-tuned LLaMA model, then concatenating the query with one of the sampled reasoning paths using a template, and finally pairing with the ground-true reasoning path to form a training sample. We use a sampling temperature of 0.7 for generating reasoning paths. During inference, we use the fine-tuned revising model to revise the prediction from the normally fine-tuned model.
The results are shown in the middle subfigure of Figure 8. We can see that with the revising model improves the final accuracy marginally comparing 36.09% to 35.90%. Surprisingly, as we increase , the performances degrade. The possible defect of the revising model is that generated samples on the training set for revising training suffer from a distribution discrepancy with generated samples on the test set for revising inference. The sampled reasoning paths on the training set may have a larger lexical similarity to the ground true reasoning paths compared to those on the test set. Therefore we try two different procedures to alleviate such an issue.
1. We use the sampled reasoning path with the largest Levenstein distance out of sampled paths with respect to the ground true path to form a training sample.
2. We split the train set to folds, and fine-tune a model on each folds and sampling reasoning path on the left fold.
The results are shown in the middle and right subfigures in Figure 8, we can see that when leveraging Levenstein distance for reasoning path selection, the fine-tuned revising model enjoys a performance boost, harvesting uniformly better performance than the fine-tuning baseline across different ’s. The results demonstrate that for the revising performance, the lexical diversity of reasoning paths matters when constructing training samples. However, the revising performance does not benefit from the -fold procedure.
Appendix E Estimating FLOPs of SFT and RFT
We mainly follow the notations of (Kaplan et al., 2020) here.
For each input sample of length in GSM8K dataset, we can split it into two parts:
where denotes the length of question and generated reasoning path and answers respectively.
where denotes the numbers of samples.
Inference FLOPs
We roughly computed the FLOPs of each token during the forward pass:
To ensure the results were more accurate and reliable, we also took into account the Key-Value (KV) cache during the decoding procedure.
Therefore, we obtain the FLOPs per token during the forward pass considering the KV cache.
The total inference FLOPs are computed as follows:
where denotes the numbers of samples. denotes the average length (tokens) of the user query and generated response respectively. In GSM8K dataset, and .