Distilling Reasoning Capabilities into Smaller Language Models

Kumar Shridhar, Alessandro Stolfo, Mrinmaya Sachan

Introduction

Large language models (LLMs) have demonstrated strong performance on a variety of reasoning tasks (Brown et al., 2020; Hoffmann et al., 2022; Chowdhery et al., 2022, inter alia). One particularly interesting strategy for prompting these models is chain-of-thought (CoT), which has been shown to elicit reasoning abilities in LLMs by asking the model to incorporate intermediate reasoning steps while solving a problem Nye et al. (2021); Wei et al. (2022b); Wang et al. (2022). However, CoT has been shown to work primarily on models with hundreds of billions of parameters Wei et al. (2022b, a) or those tuned to a wide range of tasks Chung et al. (2022); Iyer et al. (2022).

Due to the significant computational resources or expensive API calls required to access CoT-capable LLMs, we ask whether it is possible to elicit such reasoning capabilities in smaller models.Following Li et al. (2022), we argue that small and large models are relative terms and context-dependent. We consider models with billions of parameters to be large, and models with millions of parameters to be small.

Small-sized, non-fine-tuned language models are known to be poor reasoners Stolfo et al. (2022). Therefore, a possible approach to induce CoT-like reasoning abilities in smaller models would be fine-tuning them on step-by-step examples.

In our work, we propose a framework for leveraging the reasoning capabilities of LLMs to supervise the training of smaller models. This approach can be thought of as a form of knowledge distillation Hinton et al. (2015), where a larger teacher model transfers knowledge to a smaller student model. However, unlike standard knowledge distillation, our method transfers the reasoning abilities of the teacher model only using its generated solutions as a proxy, i.e., we do not assume access to the teacher model parameters. Our approach consists of prompting an LLM to produce step-by-step annotations leading to the answer for a set of problems. This annotation is then used as supervision to fine-tune the student model. A high-level illustration of the process is provided in Figure 1.

Within this framework, we study three different types of annotation structure for supervising our distillation approach: (i) We consider fine-tuning on the gold step-by-step solution procedure for datasets where the step-by-step solutions are available. (ii) We study whether procedural supervision, coming from the chain of thought (CoT) of the teacher model can improve upon the baseline. (iii) We propose a third type of supervision structure, which we call Socratic CoT. This approach relies on learning a semantic decomposition of the original problem into a sequence of subproblem-solution pairs using two models – a) a question generator that learns to decompose the problem into a sequence of subproblems, and b) a question-answering model that solves the various generated subproblems (more details are in section 3.2). This approach can be thought of as an extension of the typical chain of thought reasoning where, unlike CoT, the intermediate steps are now decomposed into subquestion-solution pairs; the subquestions guide the generation of intermediate steps that lead to the final answer to the problem.

We train distilled student models with various annotation structures mentioned above. Depending on the annotation available for the given data, we use the teacher model to generate either a CoT-like solution to a problem or, if the step-by-step annotation is available, a set of subquestions leading to the solution of the problem, or both (examples of different annotations are shown in Figure 2).

We perform our analyses on three multi-step reasoning datasets: GSM8K Cobbe et al. (2021), StrategyQA Geva et al. (2021), and SVAMP Patel et al. (2021). We consider data with various types of annotation to cover a range of realistic data scenarios. Our results show that supervision by CoT-decomposed examples helps smaller models perform better, and subquestioning introduced by Socratic CoT can provide further improvement. We observe performance gains of up to 40% with LLM-generated step-by-step annotations – this validates the effectiveness of our distillation framework (detailed analysis in Section 5).

Related Work

Solving multi-step reasoning tasks like MWPs has been a popular area of research for the last couple of years Kushman et al. (2014); Hosseini et al. (2014); Roy et al. (2015); Amini et al. (2019); Zhang et al. (2020); Shridhar et al. (2022); Opedal et al. (2023). However, the majority of the modern approaches for these problems are shifting towards using large language models, often relying on approaches involving prompting or in-context learning (Cobbe et al., 2021; Kojima et al., 2022; Wei et al., 2022b; Chowdhery et al., 2022; Lewkowycz et al., 2022; Srivastava et al., 2022). One such prompting approach is the chain of thought prompting Wei et al. (2022b), which prompts the language model to generate a series of intermediate steps that improve the reasoning capabilities in LLMs. Wang et al. (2022) took another step forward and sampled multiple reasoning paths and selected the most relevant output using majority voting. Huang et al. (2022) used the most voted outputs to further fine-tune the model for better performance. Kojima et al. (2022) further improved the reasoning of LLM in a zero-shot manner by appending “Let’s think step by step” to the prompt. In contrast, our work does not propose prompting solutions; instead, we explicitly guide the student model reasoning using sub-questions at each step. Most similar to our work is the work by Zhou et al. (2022) which decomposes questions into sub-questions and asks the language model to solve each sub-question sequentially. However, this work is also restricted to prompting and only works with LLMs with billions of parameters.

Our approach is reminiscent of knowledge distillation (Ba and Caruana, 2014; Hinton et al., 2015) in that we use a student network to mimic the large teacher language model. Snell et al. (2022) demonstrated the usefulness of providing instruction that can help models achieve better reasoning skills. Similar to our hypothesis, Eisenstein et al. (2022) argued that question-answering systems should focus not only on the final answer, but also on the rationale that justifies their reasoning, to help them reason better. We go beyond this; in our work, in addition to the question-answering system, we also focus on what questions need to be asked at each step that can help to learn that reasoning step better. Finally, similar to our hypothesis of injecting reasoning capabilities into smaller models, Li et al. (2022) used CoT-like reasoning from LLMs to train smaller models on a joint task of generating the solution and explaining the generated solution. We, on the other hand, use the LLM to generate subquestions and solution pairs and use them together to inject reasoning capabilities into smaller models.

The idea of inquiring or asking information-seeking questions for discovery learning has been studied well in the past Bruner (1961). Rao and Daumé III generated clarification questions based on Stack Exchange questions as supervision, Klein and Nabi (2019) used a joint question answering model to ask questions from a given span of text and later answer them, and Rajani et al. (2019); Shwartz et al. (2020) asked questions to improve common sense QA models. In contrast, our work focuses on multistep reasoning tasks where intermediate clarifying questions and reasoning steps may not always be available and may need to be extracted from a teacher model.

Methodology

The setting we consider consists of a data set D\mathcal{D}, where each problem PiP_{i} is accompanied by a final answer aia_{i} that can be reached by several steps of reasoning. The task of solving the problem using a model ψ\psi is to predict an answer a^=ψ(P)\hat{a}=\psi(P) such that a^=a\hat{a}=a. We consider different data scenarios where intermediate annotations of the solution may be available in different forms (e.g., step-by-step, as a semantic decomposition by subquestions) or may not be present. Depending on the availability of annotations, we propose different approaches to augment the training of a small model on D\mathcal{D} by using LLMs.

A data set may present an annotation that contains intermediate reasoning steps that lead to the answer aia_{i} (i.e., a chain-of-thought annotation). This intermediate annotation can be used directly to fine-tune a small model. However, in cases where such step-by-step information is not available, we use a LLM to generate the reasoning steps that might improve the performance of the small model.

To achieve this, we consider a small subset of the dataset D\mathcal{D} and decompose each problem PiP_{i} into nin_{i} intermediate reasoning steps. We construct these intermediate reasoning steps manually, since we only need a few examples as prompts (examples are provided in Appendix Table 6).

For each remaining problem P∈DP\in\mathcal{D}, we then prompt a large language model M\mathcal{M} to generate the intermediate reasoning steps. We make sure that the chain of reasoning steps is meaningful by checking whether the last solution matches the ground truth answer, i.e. whether ai(ni)=aia_{i}^{(n_{i})}=a_{i}, where ai(ni)a_{i}^{(n_{i})} represents the answer corresponding to the last reasoning step. If this is not the case, we discard the problem and sample a new chain by prompting the model again (for a maximum of 3 times). In this way, we obtain an augmented dataset D∗\mathcal{D}^{*} in which a subset of problems is paired with a sequence of reasoning steps leading to the correct result. Finally, we can distill the reasoning capabilities into smaller models by fine-tuning them with the generated intermediate steps.

2 Distilling step-by-step reasoning through Socratic CoT

In this section, we describe how CoT can be enhanced through subquestioning. An illustration of our approach is shown in Figure 3.

In Section 3.1, we detailed how an LLM can be used to generate the intermediate annotation of a problem PiP_{i} as a chain of steps leading to the answer aia_{i}. We now extend this procedure to include a subquestion at each step of the solution. Following a similar procedure as described in Section 3.1, we prompt the LLM with few exemplars of problems decomposed as a set of intermediate subquestion-solution pairs (the prompts are reported in Appendix Table 6). This way, we obtain an intermediate annotation that includes subquestioning. In particular, each of the nin_{i} steps constituting the overall solution is a subquestion-solution pair, denoted qi(j),si(j)q_{i}^{(j)},s_{i}^{(j)}, j∈{1,…,ni}j\in\{1,\dots,n_{i}\} (an example is shown in Figure 2). We refer to the ordered list of subquestion-solution pairs for problem PiP_{i} as (qi(1),si(1)),…,(qi(ni),si(ni))(q_{i}^{(1)},s_{i}^{(1)}),\dots,(q_{i}^{(n_{i})},s_{i}^{(n_{i})}).

2.2 Transferring the Reasoning Capability into the Student

We present two strategies to distill the reasoning annotation provided by the LLM into smaller models.

In the first strategy, a single unified student is trained to generate the subquestion-solution pairs simultaneously, while in the second strategy, the question generation and question-answering tasks are assigned to two separate models. We call this second strategy iterative because the question-answering model is trained to solve each subquestion iteratively.

Using the problems in D\mathcal{D} that contain the chain of intermediate questions and solutions, we train a unified student model Muni\mathcal{M}_{uni} that learns to generate the sequence of subquestion-solution pairs {(q(1),s(1)),(q(2),s(2)),… }\{(q^{(1)},s^{(1)}),(q^{(2)},s^{(2)}),\dots\} that lead to the solution of a given problem. We use a pre-trained transformer-based model Vaswani et al. (2017) and train it on the chain of subquestion-solution pairs for each problem PP. Given a step jj of problem PP (i.e., the concatenation of q(j)q^{(j)} and s(j)s^{(j)}) consisting of a sequence of mjm_{j} tokens {xj(1),…,xj(mj)}\{x_{j}^{(1)},\dots,x_{j}^{(m_{j})}\}, we use a typical auto-regressive language modeling loss, L\mathcal{L}:

The iterative version of the student separates the tasks of generating the subquestions and providing an intermediate answer to each subquestion into two distinct models: a question generation (QG) model and a question answering (QA) model. Both the QG and QA models are implemented using a Transformer-based language model Vaswani et al. (2017). In particular, the QA model Mqa\mathcal{M}_{qa} is iteratively trained to answer the teacher-generated sub-questions. The learning objective is computed at the token level for each intermediate solution:

where ljl_{j} and the yjy_{j}’s represent, respectively, the length and the tokens of the intermediate solution s(j)s^{(j)}. s:(j−1)s^{:(j-1)} consists of the previous solution generated by the QA model iteratively in the past iterations.

Similarly, the QG model is trained to acquire the ability of the teacher model to decompose the problem’s main question into a series of sub-steps, each of which corresponds to a subquestion. The loss for this model is analogous to Equation 1, with the only difference being that the intermediate solutions are not considered for the QG model. During training, the previous intermediate solutions generated by the QA model are replaced with the teacher-generated solutions using teacher forcing Cho et al. (2014). However, the intermediate solutions generated by the model are used at inference time.

Given an unseen problem PP, the unified student model can directly predict a solution as a sequence of subquestions and answers. In the iterative approach, we first generate the subquestions conditioning the generation of the QG model on PP. After these questions are generated, they are provided to the QA model one by one, decoding the intermediate solution s^(j)\hat{s}^{(j)} at step jj token by token according to the model’s probability distribution over its vocabulary:

where yj(k)y_{j}^{(k)} is the kk-th token being decoded in greedy fashion.

After the last solution s^(n)\hat{s}^{(n)} has been generated, the numerical prediction a^(n)\hat{a}^{(n)} is parsed from the text using simple heuristics.

We study how smaller models can learn to reason better on three multi-step reasoning datasets: GSM8K Cobbe et al. (2021), StrategyQA Geva et al. (2021), and SVAMP Patel et al. (2021). GSM8K consists of 8.5K grade school math word problems, each requiring 2 to 8 steps of reasoning to solve. The solutions primarily involve a sequence of elementary calculations using basic arithmetic operations (++, −-, ×\times, ÷\div). The dataset is divided into 7.5K training problems and 1K test problems. To evaluate the model on SVAMP, we train the model on 761 multi-step math word problems taken from the ASDiv Miao et al. (2020) training set and evaluate it on 237 multi-step SVAMP problems. For StrategyQA, the test set with facts is not available, so we split the data into 80% training, 10% as validation data, and the last 10% as test data. We do not shuffle the data to maintain reproducibility.

2 Experimental Setup

We use three kinds of annotation, corresponding to the three datasets that we consider.

: The GSM8K dataset falls into this category and includes a Socratic version where intermediate subquestion-solution pairs are provided for each MWP. While the intermediate step-by-step solutions were manually annotated, the authors report that the subquestions were generated by prompting GPT-3. We reproduced a subset of these subquestions using a GPT-3 model with prompts, and we observed a high similarity between the questions provided and the ones generated by us (BERT F1F_{1} score of 95%). For Socratic CoT, we thus use the subquestioning annotation already provided.

: We study the StrategyQA dataset, which falls in this category. Strategy QA consists of a factual question with binary True/False as the final answer. Additional supporting facts and decomposed questions are provided. However, the set of facts and the decomposed questions provided with a given question are not always aligned (i.e., a fact is not necessarily the answer to one subquestion). Therefore, having a setup similar to the one for GSM8K is not possible. We thus consider two versions of the data. One in which the supporting facts are used as CoT and the corresponding questions are generated by prompting a GPT-3 model, and a second in which we take the provided questions and generate the facts (this time aligned with the questions) using GPT-3.

: AsDiv/SVAMP falls in this category and for training, we use GPT-3 to generate both intermediate subquestions and solutions. Intermediate solutions are used as CoT and the generated subquestion-solution pairs for Socratic CoT.

3 Implementation Details

We use GPT-2 variants Radford et al. (2019) as student models. GPT-3 175B Brown et al. (2020) served as the teacher model for decomposing complex problems into a series of simpler substeps (we report the prompts used in Appendix Table 6).

All models were trained using the Huggingface library Wolf et al. (2020) on an NVIDIA Tesla A100 GPU with 40 GB of memory. Each experiment was run for the same number of iterations to ensure fairness with periodic evaluation over the validation set. Teacher forcing was used during training to replace the generated responses with ground truth answers from the training dataset.

To evaluate the question-answering performance on the GSM8K, SVAMP, and StrategyQA datasets, we compute the accuracy based on the final answer provided by the student model.

Table 4.3 demonstrates that leveraging LLMs reasoning capabilities using our framework can improve the reasoning results for all dataset types.

When human-annotated step-by-step solutions are available, training smaller models with LLM-generated CoT is not advantageous, as shown on GSM8K. This is to be expected since the annotation generated by an LLM is likely to be noisier and of lower quality than human-annotated data. However, the ground-truth step-by-step annotation can be leveraged to prompt an LLM to generate subquestions for the Socratic CoT approach, giving a performance boost of up to 38% when the LLM-generated subquestions are used at inference time. When the subquestions are learned by the QG model (Iterative SocCoT), the accuracy of the student model decreases slightly but still improves over the step-by-step annotation without subquestions (17.89 vs. 14.10). Figure 5 shows a comparison of predictions generated by SocCoT models and a model trained on the GT step-by-step annotation. Unified Socratic CoT performs similarly to training with the step-wise ground-truth annotation. We additionally include the score produced by GTP-3 6B to show that training with Socratic CoT can help a small model (GPT-2 large with 774M parameters) perform as well as a nearly 10x larger model fine-tuned with human annotated data.

On StrategyQA, we observe that the inclusion of ground-truth supporting facts in the fine-tuning procedure improves the performance of the small models. However, surprisingly, when the supporting facts are generated by GPT-3, their inclusion actually hurts performance (58.07 vs 60.51 for GPT-2 Large). We hypothesize that this is likely due to the imperfect factual knowledge provided by the LLM, which mars the quality of the supervision. We have observed that the GT supporting facts provided often do not represent a logical sequence of propositions leading to the final answer. This is likely the reason why decomposing the problem through subquestions based on such facts actually harms accuracy (see SocCoT column in Table 4.3). Instead, using the provided subquestions and using an LLM to generate the answers (representing coherent facts leading to the final answer) proves to be an effective strategy (60.31 vs. 52.02 for GPT-2 Medium). A more detailed comparison between our proposed approaches is presented in Figure 4. However, GPT-2 XL models perform well when trained on facts as unlike smaller models, larger models can encode more facts at once in their parameters, which assists in answering a factual question.

On the SVAMP dataset, which includes only final answers and no intermediate annotation, LLMs can be used to generate both the intermediate steps and the subquestions. Both the consideration of intermediate solutions without subquestions (CoT) and the consideration of intermediate solutions with subquestions (SocCoT) lead to an improvement in performance. The trend here is similar to what was observed for StrategyQA, with Socratic CoT being more effective for the two smaller models but falling back to CoT for the larger model.

We have experimented with an alternative training solution that does not involve a question-generation model. This strategy aims to improve the supervision for fine-tuning a small model through subquestioning, but without relying on the presence of subquestions at test time. The procedure consists of training the student model to generate the entire chain of steps leading to an intermediate answer. That is, when the sub-question q(1)q^{(1)} is asked, the model is trained to generate the answer s(1)s^{(1)}, but when q(j)q^{(j)} is asked, the model is trained to generate the chain of thought reasoning {s(1),s(2),…,s(j)}\{s^{(1)},s^{(2)},\dots,s^{(j)}\} (instead of just s(j)s^{(j)}). This eliminates the need for the intermediate sub-questions at inference time, as the model is trained to implicitly decompose the main problem into smaller reasoning steps. However, this method leads to significant performance degradation (results are reported in Table 5), highlighting the need for subquestions at inference time.

Conclusion

In Figures 5 and 7, we report example outputs predicted by GPT-2 models for a set of GSM8K and SVAMP problems.

The chain-of-thought style of step-by-step reasoning has proven to be very effective for reasoning in LLMs. In this work, we propose ways to distill these reasoning capabilities into smaller models and suggest ways to further improve them by explicitly asking stepwise questions. We demonstrate the effectiveness of our proposed methodology on three popular multi-step reasoning datasets, and discuss cases where one method should be preferred over the other for different datasets.

In our work, we use only one solution from the LLM to distill information into the student model, and according to Wang et al. (2022), multiple subquestion-solution pairs can be sampled, and using majority voting, all pairs leading to the most frequent answer can be used to distill knowledge into the student models. Also, due to computational budget, we used a single prompt to compare the CoT and Socratic CoT and using more prompts (up to 8) might lead to a fairer comparison and better results Wei et al. (2022b). We leave these experiments for the future.

Although this work improves the reasoning capabilities of smaller models, the models are still not powerful enough to be used in sensitive settings such as education. We plan to release our code and model checkpoints, but the models must be used carefully by users, as many generative models, including ours, are prone to hallucination.

Alessandro Stolfo is supported by Armasuisse Science and Technology through a CYD Doctoral Fellowship.