Training Chain-of-Thought via Latent-Variable Inference

Du Phan, Matthew D. Hoffman, David Dohan, Sholto Douglas, Tuan Anh Le, Aaron Parisi, Pavel Sountsov, Charles Sutton, Sharad Vikram, Rif A. Saurous

Introduction

For many mathematical, logical, and common-sense reasoning problems, large language models solve problems more accurately when instructed to work out the answer step by step in a chain of thought or a scratchpad (Wei et al., 2022; Nye et al., 2021; Kojima et al., 2022; Rajani et al., 2019; Shwartz et al., 2020). These methods encourage the model to produce a rationale, that is, text describing a sequence of reasoning steps that leads to an answer; the motivation is that it seems to be easier for the model to generate a sequence of correct reasoning steps than to generate the final answer directly. Because of the striking performance of chain-of-thought methods, many variants have been proposed (Wang et al., 2022b; Zhou et al., 2022; Creswell et al., 2022; Ye & Durrett, 2023), but there are still many cases in which the rationales are incorrect.

One way to improve these methods is to fine-tune models to generate better rationales. If gold-standard rationales can be obtained, such as via crowdsourcing (Rajani et al., 2019) or automatically (Nye et al., 2021), then supervised methods can be applied, but obtaining this data can be difficult. An appealing alternative is to start from datasets that contain questions and correct answers only, which are more readily available, and bootstrap rationales during learning. A version of this strategy was proposed as the self-taught reasoner (STaR) (Zelikman et al., 2022), which generates proposed rationales from an LLM, and then fine-tunes on rationales that lead to the correct answer.

In this paper, we approach the problem of bootstrapping rationales from a different conceptual direction: chain-of-thought methods are probabilistic latent-variable models. The LLM defines a joint probability distribution over questions, rationales, and answers; this joint distribution implies a marginal distribution of answers given questions, averaging over all possible rationales weighted by their probability given the question. The problem of self-training for reasoning then becomes one of learning with incomplete data, a core task in probabilistic machine learning (Murphy, 2022) to which we can apply methods from a large and sophisticated literature.

This perspective raises a technical challenge, because computing the marginal distribution requires averaging over a vast set of potential rationales. To address this, we introduce a learning algorithm for rationale generation, which we call TRICE.TRICE stands for “Tuning Rationales with Independence-Chain Expectation-maximization.” TRICE is a simple Markov-chain Monte Carlo (MCMC) expectation-maximization (EM) algorithm combined with a novel control-variate scheme, inspired by ideas from STaR (Zelikman et al., 2022), memoized wake-sleep (Hewitt et al., 2020), Markovian score climbing (Naesseth et al., 2020), and persistent contrastive divergence (Tieleman, 2008).

This view unifies several threads of work in reasoning using LLMs: It provides an alternative interpretation of STaR as a kind of biased stochastic expectation-maximization algorithm (Nielsen, 2000) that underweights difficult examples when its rationalization process fails. Self-consistency (Wang et al., 2022a) can be seen as a Monte Carlo algorithm for computing the most likely answer under the marginal distribution. Compared to self-consistency, the probabilistic learning approach of TRICE allows us to average over rationales not only at inference time, but also at training time. Compared to STaR, TRICE is less likely to ignore difficult examples (which stabilizes convergence and improves performance), and is also able to learn from incorrect rationales as well as correct ones.

We apply our technique to the GSM8K dataset (Cobbe et al., 2021) and to the BIG-Bench Hard benchmark (Suzgun et al., 2022a). We find that TRICE improves the model’s performance significantly, outperforming models tuned with STaR, direct tuning with or without CoT, and even supervised fine-tuning on human-generated rationales.

Method

Given a training set of NN questions x1:Nx_{1:N} and answers y1:Ny_{1:N}, we formalize CoT tuning as optimizing a parameter vector θ\theta to maximize the average marginal log-likelihood of answers given questions:

where zz is an unobserved latent rationale, pθ(z∣x)p_{\theta}(z\mid x) is the probabilityUnless otherwise specified, we sample at temperature 1 throughout. of obtaining the rationale zz by prompting an LLM with the question xx and tunable parameters θ\theta, and pθ(y∣z,x)p_{\theta}(y\mid z,x) is the probability of obtaining the answer yy given rationale zz, question xx, and parameters θ\theta. We will be particularly interested in models where the likelihood pθ(y∣x,z)∈{0,1}p_{\theta}(y\mid x,z)\in\{0,1\}, that is, where the answer yy is a deterministic function of zz. For example, we might say that the model’s answer is y=“(a)”y=\textrm{``(a)''} if zz ends with the string "The answer is (a)." For this deterministic model, we define p(y∣z,x)=c(z,y)∈{0,1}p(y\mid z,x)=c(z,y)\in\{0,1\}. Details of c(z,y)c(z,y) for each task can be found in Appendix F. We believe that such a binary likelihood model is appropriate for question-answering tasks where zz is a rationale—a good rationale should leave no ambiguity about the correct answer. The derivations below will therefore assume a binary likelihood function. It is straightforward to generalize our methods to cases where the relationship between zz and yy is weaker and therefore p(y∣x,z)p(y\mid x,z) is more complicated; Appendix A shows how.

Algorithm 1 outlines the method. A notebook with a reference implementation can be found at https://github.com/google-research/cascades/tree/main/cascades/examples/notebooks/trice.ipynb.

We start by initializing a memory containing a latent rationale znz_{n} for each example pair xnx_{n}, yny_{n} by sampling znz_{n} from a hinted guide distribution q(z∣xn,yn)q(z\mid x_{n},y_{n}) that may condition on the correct answer yny_{n} as well as the question xnx_{n}. For example, the guide might prompt an LLM specifically to give an rationale for the answer; more details about the precise prompts used by the guide are in Appendix F. In some cases sampling from the guide instead of the model pθ(z∣xn)p_{\theta}(z\mid x_{n}) increases the chances of generating a correct rationale (Zelikman et al., 2022).

At this point we have all we need to compute a gradient estimate; we can just average the gradients ∇θlog⁡pθ(zim∣xim)\nabla_{\theta}\log p_{\theta}(z_{i_{m}}\mid x_{i_{m}}) that we obtain from those rationales in the updated memory that are correct (i.e., we ignore examples where both the proposed rationale and the previous rationale were wrong). basic_gradient_estimate in Algorithm 1 shows how.

control_variate_gradient_estimate is more expensive than basic_gradient_estimate, since we must compute gradients not only for the rationales in memory but also for any incorrect rationales we generate. This may be wasteful, especially if many of the weights on those gradients (1−β1-\beta for correct proposals, β\beta for incorrect proposals) are close to zero because β\beta is close to zero or one. To reduce this cost, in subsampled_control_variate_gradient_estimate, we use systematic resampling (Hol et al., 2006) to generate a subsample of LL question-rationale pairs, from which we obtain an unbiased estimate of the output of control_variate_gradient_estimate. We preferentially sample gradients with higher scalar weights; if β\beta is small, we are less likely to sample incorrect rationales (which have weight β\beta), and if β\beta is large, we are less likely to sample correct proposed rationales (which have weight 1−β1-\beta). This can be seen as a generalization of the strategy of Burda et al. (2015, Section 3) for reducing the cost of computing IWAE gradients.

Below, we derive this variance-reduced stochastic MCMC-EM procedure in more detail.

The gradient of the marginal log-likelihood log⁡pθ(y∣x)\log p_{\theta}(y\mid x) with respect to θ\theta is

that is, it is the expectation with respect to the posterior pθ(z∣x,y)p_{\theta}(z\mid x,y) of the gradient of the conditional log-prior log⁡pθ(z∣x)\log p_{\theta}(z\mid x), since the likelihood p(y∣z,x)=c(z,y)p(y\mid z,x)=c(z,y) does not depend on θ\theta. So if we can sample from the posterior over rationales zz conditioned on the question-answer pair x,yx,y, then we can compute an unbiased estimate of the gradient of the marginal log-likelihood log⁡pθ(y∣x)\log p_{\theta}(y\mid x). We can interpret this as “bootstrapping” rationales zz that are consistent with both the prior on rationales pθ(z∣x)p_{\theta}(z\mid x) and the observed answer yy (cf. Zelikman et al., 2022).

We cannot directly sample from pθ(z∣x,y)p_{\theta}(z\mid x,y), so we resort to Markov chain Monte Carlo (MCMC). We maintain a memory (cf. Hewitt et al., 2020) of a single rationale znz_{n} for each question-answer pair xn,ynx_{n},y_{n}, and each iteration we apply a random update to znz_{n} that leaves the posterior pθ(zn∣xn,yn)p_{\theta}(z_{n}\mid x_{n},y_{n}) invariant (cf. Tieleman, 2008). Each MCMC update brings the znz_{n}’s closer in distribution to pθ(zn∣xn,yn)p_{\theta}(z_{n}\mid x_{n},y_{n}) (Cover, 1999; Murray & Salakhutdinov, 2008). However, updates to θ\theta may change the posterior pθ(zn∣xn,yn)p_{\theta}(z_{n}\mid x_{n},y_{n}), so we must keep updating the chains to control the bias of our gradient estimates.

Remarks: Independence samplers can be understood as “Metropolized” importance samplers that spread the work of generating and evaluating proposals over time. In our setting, the update can also be interpreted as attempting to sample from the posterior by rejection sampling, then falling back on an old sample if that fails. The expected number of iterations between successful updates is p(y∣x)−1p(y\mid x)^{-1}, so mixing will be faster for easier questions xx, and will accelerate as the model improves.

Basic gradient estimator.

Remarks: The estimate will have low bias if the distribution of z′z^{\prime} is close to the posterior p(z∣x,y)p(z\mid x,y), which we expect to be true if the chain is mixing quickly enough relative to how fast θ\theta is changing. This will happen if either the probability of getting a correct answer is high, or if θ\theta is changing slowly due to a small learning rate and/or gradient. If the model’s training-set accuracy improves with training and we use a decaying learning-rate schedule, then as training proceeds both of these factors should work to reduce the bias of the gradient estimate.

Adding a control variate.

Estimating β𝛽\beta.

Gradient subsampling.

2 Why not variational inference, reweighted wake-sleep, or rejection sampling?

We considered three alternatives to the MCMC-EM approach that we pursue in this paper: variational EM (e.g., Bishop, 2006), reweighted wake-sleep (RWS; Bornschein & Bengio, 2015; Le et al., 2019), and rejection sampling.

Variational expectation-maximization is a common strategy for training latent-variable models, but variational inference with discrete latent variables is challenging (e.g., Tucker et al., 2017).

RWS is an attractive alternative that avoids high-variance score-function gradients; it proceeds by sampling MM samples z1:Mz_{1:M} from a guide model qϕ(z∣x,y)q_{\phi}(z\mid x,y), assigning the samples weights wm∝pθ(y,z∣x)qϕ(z∣x,y)w_{m}\propto\frac{p_{\theta}(y,z\mid x)}{q_{\phi}(z\mid x,y)}, and updating both the model parameters θ\theta and the guide parameters ϕ\phi to maximize the reweighted log-probabilities ∑mwmlog⁡pθ(zm∣x)\sum_{m}w_{m}\log p_{\theta}(z_{m}\mid x) and ∑mwmlog⁡qϕ(zm∣x,y)\sum_{m}w_{m}\log q_{\phi}(z_{m}\mid x,y). Unfortunately, we found that RWS training sometimes led to degenerate zero-length rationales zz. Figure 1 suggests a partial explanation: shorter sequences get higher weights, so the model and guide learn to produce shorter and shorter sequences until they consistently produce empty rationales.

With careful initialization and learning-rate tuning, we could sometimes get RWS to avoid this problem of empty rationales. But this led to a new problem: the guide qϕ(z∣x,y)q_{\phi}(z\mid x,y) learned to closely mimic the prior p(z∣x)p(z\mid x) until the very end of the rationale, and then simply paste in the correct answer whether or not it had anything to do with the rationale up to that point (cf. Turpin et al., 2023). Figure 5 in Appendix E shows a representative example in which the guide model ignores the answer it arrived at through incorrect reasoning and pastes in the correct answer.

Quantitatively, denoting by tt the index of the token at which the “final answer” section of the rationale begins, in one run we found that the average KL between q(z1:t∣x,y)q(z_{1:t}\mid x,y) and p(z1:t∣x)p(z_{1:t}\mid x) was about 0.610.61 nats, while the conditional KL between q(z(t+1):T∣x,y,z1:t)q(z_{(t+1):T}\mid x,y,z_{1:t}) and p(z(t+1):T∣x,z1:t)p(z_{(t+1):T}\mid x,z_{1:t}) was about 42.542.5 nats, confirming that the guide was not “reasoning backwards”, just copying the correct answer.

Finally, we considered a rejection-samplingWe also considered optimizing an importance-weighted bound (Burda et al., 2015) using the prior p(z∣x)p(z\mid x) as a proposal distribution, but instead opted for a simple rejection sampling scheme since this is less biased and equally feasible in our setting. scheme in which we sample KK proposal rationales z1:Kz_{1:K} from p(z∣x)p(z\mid x), and average the gradients from those rationales that lead to correct answers. We will present the quantitative results in Section 4; our main finding is that, while this scheme can work, it requires reducing the minibatch size by a factor of KK to keep the per-iteration cost constant compared to TRICE, which in turn leads to slower convergence and/or worse final results.

Related Work

A number of methods have proposed rationale generation for problem-solving tasks in neural sequence models, including both fully supervised and few-shot approaches (Wei et al., 2022; Nye et al., 2021; Kojima et al., 2022; Rajani et al., 2019; Shwartz et al., 2020; Wang et al., 2022b; Zhou et al., 2022; Creswell et al., 2022; Ye & Durrett, 2023). Particularly relevant to our approach is self-consistent chain-of-thought (Wang et al., 2022b), because this can be approximately viewed as marginalizing over rationales at test time. This technique has been successfully applied for a range of quantitative reasoning tasks (Lewkowycz et al., 2022). There is relatively much less work that does imputation or averaging over rationales at training time; perhaps the main instance is STaR (Zelikman et al., 2022), which we discuss in Section 3.1.

Dohan et al. (2022) present a position paper which advocates representing a composition of language model interactions via probabilistic programming. Our treatment of rationales as latent variables is inspired by that work. Lievin (2022) offers another example of interpreting LLMs with CoT as latent-variable models.

Variational inference (e.g., Kingma & Welling, 2013) and wake-sleep methods (e.g., Bornschein & Bengio, 2015) are workhorses of the latent-variable-modeling community, but as we discuss in Section 2.2 we found the bias of these methods to cause serious problems. MCMC-EM is a less-common strategy these days, although a version of it based on Gibbs sampling (Geman & Geman, 1984) it has been widely applied to training undirected graphical models (Tieleman, 2008). TRICE can also be cast as an instance of Markovian score climbing (Naesseth et al., 2020).

ReAct (Yao et al., 2023) demonstrated that injecting reasoning into an RL-style observe-and-act loop significantly increases performance. This approach was extended in Reflexion (Shinn et al., 2023), where an agent can conditionally reflect on an RL trajectory, augmenting the resulting examples which can be used as few-shot examples in subsequent rollouts. These approaches reported significant improvements on their respective evaluation tasks but still rely on the model being able to produce useful and actionable feedback through pure few-shot prompting, whereas our method actively tunes the model to produce thoughts amenable to the task.

Recent work on tool-use within language models also works via imputation, inferring where to insert calls to tools (Parisi et al., 2022; Schick et al., 2023). Their loss functions are similar in spirit to ours, filtering out trajectories which do not lead to valid answers. In this paper, we have treated rationales as latent variables; one could also treat tool-use as a latent variable.

The most closely related work is the self-taught reasoner (STaR; Zelikman et al., 2022). Besides the arguments in their derivations, there are three significant differences between TRICE and STaR. First, STaR uses greedy decoding, which reduces the diversity of the rationales it trains on. The authors made this choice to reduce the danger of the model getting the right answer despite having a bad rationale. While we do find that our procedure sometimes generates correct answers for the wrong reasons, this did not seem to stand in the way of the model improving on most tasks. One reason may be that our base models are more powerful than the 6B-parameter GPT-J model used in the STaR paper, so they are more likely to generate good rationales from the beginning.

A second difference is that TRICE resamples rationales every iteration, so it are less likely to overfit to any particular rationale. STaR has an inner loop that runs many training iterations on a single set of rationales, meaning it uses stale rationales to estimate the gradient of the marginal likelihood. In our experiments, we observed that this leads to the model effectively memorizing a fixed set of rationales for the training set. Once this happens, the greedy decoding procedure will almost certainly reproduce exactly the same rationales at the beginning of the next outer loop. If these rationales all lead to the correct answer, and STaR has a rationale for each question, then this is a global optimum of the marginal likelihood on the training set! But empirically, STaR often does not find a good rationale for each question, and so it ignores some fraction of the training set (see Section 4).

A final, minor difference is that when STaR updates its rationales, it may replace a rationale from the model p(z∣x)p(z\mid x) with a rationale from a surrogate qθ(z∣x,y)q_{\theta}(z\mid x,y). As the model memorizes a set of correct rationales for the training set, STaR becomes less likely to fall back on the surrogate, but this choice could affect early training dynamics.

Experiments

We evaluate TRICE on the GSM8K (Cobbe et al., 2021) dataset and the 27 BigBench-Hard (BBH) tasks (Suzgun et al., 2022b) using the medium-size PaLM 2-M (Anil et al., 2023) Transformer-based LLM (Vaswani et al., 2017). For the BBH experiments, we used the Flan instruction-tuned (Chung et al., 2022) version of PaLM 2; for GSM8K, we used the base PaLM 2 model, since GSM8K is included in the Flan training datasets. All experiments were run on TPU v4 and v5e chips (Jouppi et al., 2023). Examples of generated rationales can be found in Appendix E.

Rather than fine-tune the model weights, we use prompt tuning (Lester et al., 2021); we prepend a sequence of embedding vectors θ\theta (a “soft prompt”) to the embeddings corresponding to the tokenized CoT prompt used to condition the model. Prompt tuning can achieve similar accuracy gains to full fine-tuning, but using a small fraction of the parameters. We initialize the soft prompt with the embedding sequence obtained from a series of three (for BBH) or five (for GSM8K) exemplar CoT prompts, each of the form “Question: \nAnswer: Let’s think step by step.\n”. We consider two initialization schemes: one where we use the standard few-shot CoT prompts that are provided with BBH, and one where we try to bootstrap a few-shot CoT prompt by sampling random questions from the training set, generating random rationales from the base model, and picking three or five examples where the random rationales lead to correct answers. The first scheme can be seen as a way of fine-tuning a good initial few-shot prompt, but it does require a small amount of detailed CoT supervision, while the second scheme only requires label supervision.

On each BBH task, we split the examples into 6060% train and 4040% test sets. For all but three tasks, this is 150150 training and 100100 test examples. For GSM8K, we use the standard 74737473-example training set and 13191319-example test set. We evaluate CoT models’ accuracy in two ways: first, using greedy (temperature-0) decoding, and second, using “self-consistency” (Wang et al., 2022b). In self-consistency evaluation, we draw 40 samples and check whether the most common answer is correct; this is a plug-in estimator for the prediction arg⁡max⁡yp(y∣x)\arg\max_{y}p(y\mid x) that minimizes 0-1 loss under the model (although this is not how Wang et al. (2022b) originally motivated the procedure).

We compare against four baseline prompt-tuning methods: direct prompt tuning, CoT prompt tuning, rejection sampling, and STaR (Zelikman et al., 2022). All methods are evaluated against the same validation sets, and use the same training labels, few-shot prompts (except for direct tuning, where we only use question-answer pairs), and initialization strategies as appropriate. Details for each method and its corresponding experimental hyperparameters can be found in Appendix F.

Section 4 and Table 2 summarize the results; more detailed task-by-task BBH summaries are in Appendix D. Even with no human-generated exemplar rationales, TRICE is able to learn to generate rationales that lead to the correct answer. TRICE also outperforms a model trained directly on human-generated rationales on GSM8K (cf. Uesato et al., 2022), perhaps because the cross-entropy loss used in supervised fine-tuning may place more weight on style than substance; it takes far more bits to encode how one expresses a chain of reasoning than it does to encode the reasons themselves.

Initializing the soft prompt with a human-generated 3-shot exemplar question-rationale-answer prompt slightly improves performance on BBH, as does evaluating with self-consistency. By the end of training, TRICE has managed to generate at least one valid rationale for almost all training examples, while STaR fails to generate valid rationales for almost 10% of training examples. Unlike in the experiments done on Commonsense QA (Talmor et al., 2019) by Zelikman et al. (2022), STaR does not outperform the direct-prompted prompt-tuned model on BBH. This may be because each BBH task includes relatively little training data (150 examples as opposed to CommonsenseQA’s 9,741), and so in its inner loop STaR overfits to its relatively small set of bootstrapped rationales. TRICE, on the other hand, can overfit to the small set of questions but at least has a chance to generate a somewhat diverse set of rationales from those questions.

One piece of evidence for this overfitting-rationales hypothesis is that on the final step of its final inner loop, STaR (with bootstrapped initialization) achieves a training sequence-level (not per-token) cross-entropy loss of less than 0.06 on all tasks, and of less than 0.01 on 19 out of 27 tasks. This implies that it has learned to exactly reproduce a single set of rationales with very high probability, which makes it very likely that it will generate those same rationales in the next iteration.

Figure 2 compares estimates for GSM8K of the average training marginal likelihood (i.e., how often a proposal is accepted) and the validation accuracy with greedy decoding as a function of number of training stepsWe set the cost per iteration of rejection sampling and TRICE with and without the control-variate scheme to be directly comparable: for rejection sampling, we reduce the minibatch size by a factor of four and generate four times as many proposals per example; for TRICE with the control-variate scheme, we set the gradient minibatch size LL equal to the number of examples per minibatch MM (note that this does still involve subsampling, since each example could potentially contribute both a correct and an incorrect rationale to the gradient estimate). for rejection sampling and for TRICE with and without the control-variate scheme. The control-variate scheme improves average convergence speed, particularly towards the end of training as the probability of generating correct answers on the training set increases. Both versions of TRICE converge to high training accuracy much faster than rejection sampling.

We proposed TRICE, a method for tuning LLMs to be better at solving question-answering tasks using chain-of-thought (CoT) prompting. By framing the CoT-prompted LLM as a latent-variable model, we were able to derive a principled and effective fine-tuning method. When applied to GSM8K and BIG-Bench Hard (BBH) tasks, TRICE outperforms three strong baselines: direct prompt-tuning, STaR, and rejection sampling. While we derived TRICE in the context of CoT question-answering, its basic MCMC-EM strategy could be employed more broadly, for example to tool-use problems.

We only evaluated TRICE with prompt-tuning on a medium-size LLM; it may be that it behaves differently on smaller models, larger models, or when using other fine-tuning strategies. TRICE is a gradient-based tuning algorithm, but many of the most capable LLMs are proprietary, and their owners often do not provide any public mechanism for gradient-based fine-tuning. This makes it hard to evaluate how well TRICE would work when used with, say, GPT-4 (OpenAI, 2023). Finally, our quantitative evaluations focused on whether the LLM could produce the right answer; we did not formally evaluate the quality of the reasoning in the rationales themselves (cf. Uesato et al., 2022).

Broader Impacts:

This work aims to improve the capabilities of LLMs by making them better able to answer questions accurately and transparently. However, more-capable LLMs may be used in malicious or unsafe ways, fine-tuning on uncurated question-answering datasets may introduce biases into the models, and more widely used LLMs will contribute a larger carbon footprint.

Rationales may make it easier for motivated users to judge the trustworthiness of LLM outputs. But many users may not read and critique an LLM’s rationales, taking the mere existence of a rationale as evidence of truth. If chain-of-thought rationales promote uncritical trust, they could lead to harm.

Acknowledgements:

We appreciate Daniel Freeman and Enrique Piqueras’ contributions to the infrastructure that we used in our experiments. We thank Kevin Murphy, Ben Lee, Brian Patton, and Jascha Sohl-Dickstein for helpful discussions.

Appendix A Generalizing TRICE to Nondeterministic Likelihood Models

To apply TRICE beyond question-answering problems, we might want to use a nondeterministic likelihood model. For example, our desired output yy might be a summary of a text xx, and zz might be a scratchpad or outline. In situations like this, there might be many yy’s that are appropriate for a given xx and zz. So from a modeling perspective, it could make sense to make p(y∣x,z)p(y\mid x,z) have nonzero entropy. But there is also a computational reason to prefer such a model: as the number of reasonable values that yy could take given xx increases, the probability of sampling the precise yy that was observed goes down at a rate that might be exponential in the size of the output space.

Fortunately, we can easily extend TRICE to accommodate nondeterministic likelihoods. The differences are:

The value of the control variate in this setting may be less than it is in the deterministic-likelihood setting. Even if we learn a model that consistently produces good latents zz (in the sense that they lead to valid outputs yy), this does not guarantee that it will consistently generate latents that are consistent with the particular yy that was observed. For example, there might be multiple reasonable ways to outline a long text, some of which lead to different summary paragraphs. In this scenario, the acceptance probability may not converge to something close to 1, and the variance-reduction effect from the control variate will be modest.

Appendix B Derivation of the Control Variate Scaling Heuristic

So the variance of g^\hat{g} simplifies to

Taking the derivative with respect to β\beta shows that this is minimized when

Plugging this back into Equation 12 gives the optimal variance v⋆v^{\star}:

where in the second line we again use the fact that πg+=−(1−π)g−\pi g_{+}=-(1-\pi)g_{-}, and in the third-to-last line we approximate 1π\frac{1}{\pi} with the first-order Taylor approximation 1π=2−π+O((1−π)2)\frac{1}{\pi}=2-\pi+O((1-\pi)^{2}). Thus, we can write

By contrast, plugging our heuristic value of β=π\beta=\pi into Equation 12 gives the suboptimal variance vπv^{\pi}:

where we use the approximation πk=(1−(1−π))k=1−k(1−π)+O((1−π)2)\pi^{k}=(1-(1-\pi))^{k}=1-k(1-\pi)+O((1-\pi)^{2}) to simplify the π2\pi^{2} and π3\pi^{3} terms. Thus, we conclude that v⋆v^{\star} and vπv^{\pi} are the same up to O((1−π)2)O((1-\pi)^{2}), and so as the probability π\pi of getting the correct answer increases, the suboptimality of setting β=π\beta=\pi goes down faster than the variance does.

Appendix C On Gradient Estimators Based Solely on Incorrect Rationales

We adopt the same shorthands as in Appendix B.

which relates the gradient we want to estimate g+g_{+} (the expected gradient given that the rationale is correct) to g−g_{-} (the expected gradient given that the rationale is incorrect).

Deferring for the moment the difficulty in estimating π−1\pi^{-1} (see Section C.2 below), we can consider the variance of an estimator based on the right hand side of Equation 20:

so that unless the variance v−v_{-} of incorrect rationales is very low, the variance of this estimator will be O(π−2)O(\pi^{-2}), which is very high. By contrast, the variance of a gradient estimator based purely on correct rationales is simply v+v_{+}, so unless the gradient variance for incorrect rationales is dramatically lower than that for correct rationales, then if π\pi is small then incorrect rationales will lead to much noisier gradient estimates.

On the other hand, if 1−π1-\pi is small, then we have

which is likely a significant improvement on the variance v+v_{+} of the correct-rationale estimator; in particular, it goes to zero as π\pi approaches 1.

C.2 TRICE control variate as a debiased estimator based on incorrect rationales

Instead, we can compute a biased estimator that ignores the π−1\pi^{-1} term and then correct for the bias:

Appendix D BBH Per-Task Experimental Results

Table 3 summarizes our experimental results for each task in BBH.

Appendix E Example Rationales

Figure 3 illustrates some examples of rationales generated by the TRICE-tuned LLM on the BBH task Logical Deduction Three Objects.

Figure 4 illustrates some examples of rationales generated by the TRICE-tuned LLM on GSM8K. Although we did find examples where the LLM got the answer right for the wrong reasons, this was much less common on GSM8K than on BBH, since the numeric output space for GSM8K is much larger than that for the typical multiple-choice BBH task.

Figure 5 shows an example where the guide model qϕ(z∣x,y)q_{\phi}(z\mid x,y) in reweighted wake-sleep learns to closely mimic the prior model pθ(z∣x)p_{\theta}(z\mid x) until the very end of the rationale, at which point it pastes in the correct answer.

Appendix F Method and Template Details

In this section, we present more details on the methods and templates that we used in the experiments.

To sample from pθ(z∣x)p_{\theta}(z\mid x), we prompt the LLM with the template “Question: \nAnswer: Let’s think step by step.\n”. We cap the length of the generated rationales at 1.25 times the length of the longest of the exemplar rationales used to initialize the soft prompt. To initialize the memory (i.e., to sample from q(z∣x,y)q(z\mid x,y) in line 2 of Algorithm 1), on BBH we sample from the base model with a “guide” prompt of the form “Question: \nAnswer: The answer is . Let’s think step by step.\n”. We use the same guide prompt to generate rationalizations in STaR, but with temperature 0 (see below). On GSM8K, we instead initialize the memory with samples from pθ(z∣x)p_{\theta}(z\mid x), since we found that initializing the memory using a prompt that includes the answer led to slower convergence and worse results; it may be that including the answer in the prompt sometimes encourages the model to produce untrustworthy explanations (Turpin et al., 2023).

To evaluate the correctness c(z,y)c(z,y) of a rationale zz given the answer yy, in BBH we search the end of the rationale for final answers in the form “the answer is .”. In GSM8K, we initialize the soft prompt to encourage the model to wrap its answers in “” and “” tags, and then extract the answer from those tags. To encourage the bootstrapped few-shot examples in GSM8K to follow this template, we the following example to the CoT prompt: “Question: What is 1 plus 1?\nAnswer: Let’s think step by step.\n1 plus 1 is 2.\n\n2\n\n\n”.

Figure 6 and Figure 7 show the bootstrapped few-shot CoT examples that we used to initialize the soft prompt in the experiments.

For all BBH tasks, we run TRICE for 500500 steps with batch size M=32M=32 and do not use subsampling (i.e., compute L=64L=64 gradients per batch). We use the Adam optimizer (Kingma & Ba, 2015) with an initial learning rate 1.01.0 and a cosine decay schedule (Loshchilov & Hutter, 2017) that reduces the learning rate by 10x over the first 450450 steps. For GSM8K, we run TRICE for 50005000 steps with a constant learning rate of 1.0, batch size M=128M=128, and compute L=128L=128 gradients per batch.

STaR.

We use an adaptation of the STaR strategy proposed by Zelikman et al. (2022), where we do prompt-tuning rather than fine-tuning on all weights. The method alternates between updating its memory and retuning the model from scratch on the updated memory in an inner loop. We apply this procedure for 10 outer-loop steps. Following Zelikman et al. (2022), we start with 4040 inner-loop optimization steps, increasing the number of inner-loop steps by 2020% each outer-loop iteration up to a maximum of 200 steps. If the training loss goes below 0.01 we break out of the inner loop. For STaR’s inner-loop optimization, we use the same prompt-tuning initialization, Adam hyperparameters as above, but with cosine decay from 1.0 to 0.1 over the course of each inner loop. To update the STaR memory, we first try generating a rationale from the prompt-tuned model by greedy decoding, then if that rationale is incorrect fall back on a rationalization generated by greedy decoding from the same guide model we use in TRICE to initialize the MCMC memory, and finally if neither procedure generates a valid rationale we omit the example from the memory.

Rejection Sampling.

We reduce mini-batch size by 44 and draw 44 rationales for each example in the mini-batch. We use the same mini-batch size, train steps, and optimizer as in TRICE for all BBH and GSM8K experiments. In BBH, we use the initial learning rate 1.0 as in TRICE. In GSM8K, we use the learning rate 0.10.1 because it achieved better results than learning rate 0.30.3, and the training procedure became unstable with learning rate 1.01.0.

CoT Prompt Tuning.

To do supervised CoT tuning, we prompt-tune the model to maximize the log-likelihoods of the training rationales given questions. The BBH datasets include very few exemplar rationales, so we cannot apply this strategy to BBH. On GSM8K, we use the same hyperparameters as in TRICE except that we early-stop the algorithm after only 1000 train steps, since the model overfits badly when we run longer.

Direct Prompt Tuning.

In this method, the model tries to guess the answer directly without generating a rationale; prompt-tuning to maximize the log-likelihood of the answers in this setup is straightforward, since there is no latent rationale to integrate out. We initialize the soft prompt using 3 examples from the training set and truncate its length to 64. The optimization procedure is carried out over 150 steps with batch size 16 and the same Adam hyperparameters as above, except that the cosine decay period is set to 150 instead of 450. We found these adjustments to the hyperparameters from different training schemes were necessary to reduce overfitting.