MCTS is a replay buffer is a gradient estimator


In a previous post, I argued that LLM RL is weighted NLL over a replay buffer of token-level transitions, trained on pretraining infrastructure. This post takes the buffer seriously: if you branch during sampling — several actions drawn from the same state, several futures continued from the same action — the buffer is not a bag of independent paths but a tree. That is an MCTS-shaped object, and the right way to train on it falls out of gradient estimator theory.

The pitch: the policy gradient hides four expectations — over states, over actions, and over the futures inside QQ and VV. A single rollout estimates all four with the same one sample. A tree gives you many samples of each, and Rao-Blackwellization turns them into lower variance without bias — but only if each expectation is conditioned on the right part of the tree. MCTS already implements two of the four corrections: value backup is the subtree-mean QQ, and visit counts are the state weights. Gradient estimator theory supplies the other two.

The estimator

The goal is a policy π(y∣x)\pi(y\mid x) that answers question xx with response yy, maximizing

J[π]=Ex∼p(x)[Ey∼π(y∣x)[r(x,y)]].\begin{equation} J[\pi] = E_{x\sim p(x)}\left[E_{y\sim \pi(y\mid x)}[r(x,y)]\right]. \end{equation}

The policy gradient, with the autoregressive decomposition into states st=(x,y<t)s_t = (x, y_{<t}) and actions at=yta_t = y_t, is

∇J[π]=E[r(x,y) ∇log⁡π(y∣x)]=E[∑tr(x,y) ∇log⁡π(at∣st)]\begin{equation} \nabla J[\pi] = E\left[r(x,y)\,\nabla \log \pi(y\mid x)\right] = E\left[\sum_t r(x,y)\, \nabla \log \pi(a_t\mid s_t)\right] \end{equation}

(exactly, no approximation — see the previous post for why token-level is the right abstraction).

Now take the tt-th term and factor the trajectory distribution around position tt:

p(x,y)=p(st) π(at∣st) p(y>t∣st,at),\begin{equation} p(x, y) = p(s_t)\,\pi(a_t\mid s_t)\,p(y_{>t}\mid s_t, a_t), \end{equation}

where p(st)p(s_t) is the marginal over prefixes. The score ∇log⁡π(at∣st)\nabla\log\pi(a_t\mid s_t) does not depend on the future, so the future marginalizes into a value:

E[r ∇log⁡π(at∣st)]=Ep(st)[Eπ(a∣st)[Q(st,a) ∇log⁡π(a∣st)]],Q(s,a)=Ep(y>t∣s,a)[r].\begin{equation} E\left[r\,\nabla\log\pi(a_t\mid s_t)\right] = E_{p(s_t)}\left[E_{\pi(a\mid s_t)}\left[Q(s_t,a)\,\nabla\log\pi(a\mid s_t)\right]\right], \qquad Q(s,a) = E_{p(y_{>t}\mid s,a)}\left[r\right]. \end{equation}

The score has zero mean under its own distribution,

Eπ(a∣s)[∇log⁡π(a∣s)]=∑aπ(a∣s)∇π(a∣s)π(a∣s)=∇∑aπ(a∣s)=∇1=0,\begin{equation} E_{\pi(a\mid s)}\left[\nabla \log \pi(a\mid s)\right] = \sum_a \pi(a\mid s) \frac{\nabla \pi(a\mid s)}{\pi(a\mid s)} = \nabla \sum_a \pi(a\mid s) = \nabla 1 = 0, \end{equation}

so we can subtract any baseline that doesn’t depend on the action. Subtracting V(s)=Eπ(a∣s)[Q(s,a)]V(s) = E_{\pi(a\mid s)}\left[Q(s,a)\right] and summing over positions:

∇J=∑tEp(st)[Eπ(a∣st)[(Q(st,a)−V(st)) ∇log⁡π(a∣st)]].\begin{equation} \nabla J = \sum_t E_{p(s_t)}\left[E_{\pi(a\mid s_t)}\left[\big(Q(s_t,a) - V(s_t)\big)\,\nabla\log\pi(a\mid s_t)\right]\right]. \end{equation}

Count the expectations: the state marginal p(s)p(s), the action distribution π(a∣s)\pi(a\mid s), the futures inside QQ, and the actions-and-futures inside VV. Four. The single-rollout estimator collapses all four onto the same one sample: the visited prefix estimates p(s)p(s), the sampled token estimates π(a∣s)\pi(a\mid s), the rollout’s own return estimates QQ, and a group statistic estimates VV. Everything that follows is about doing better.

Rao-Blackwellizing the action expectation

Start with π(a∣s)\pi(a\mid s), because it’s the one expectation you can often compute exactly. Importance weights and sampling variance exist because an expectation is intractable. At the sequence level that’s unavoidable — exponentially many responses. But the inner expectation in equation (6) is over the vocabulary, and the forward pass already computes the full distribution πθ(⋅∣s)\pi_\theta(\cdot\mid s) at every position. For any per-token quantity g(s,a)g(s,a) that can be evaluated at arbitrary actions,

Ea∼πθ(⋅∣s)[g(s,a)]=∑a∈Vπθ(a∣s) g(s,a)\begin{equation} E_{a\sim\pi_\theta(\cdot\mid s)}\left[g(s,a)\right] = \sum_{a\in \mathcal{V}} \pi_\theta(a\mid s)\, g(s,a) \end{equation}

is an exact sum under the current policy — no action sampling, no importance ratio, zero variance from the action choice. The familiar instance is the KL penalty: summing the per-token KL to the reference policy over the vocabulary, instead of using a sampled-token estimator, is exactly this move.

When evaluating gg at all KK tokens is too expensive, Liu et al. (2019) give the right compromise. Let CkC_k be the kk highest-probability tokens under qη=πθ(⋅∣s)q_\eta = \pi_\theta(\cdot\mid s) and Cˉk\bar{C}_k the rest. Compute the head exactly, and estimate the tail with a single draw from the conditional distribution:

g^=∑z∈Ckqη(z) g(z)+qη(Cˉk) g(v),v∼qη∣Cˉk.\begin{equation} \hat{g} = \sum_{z\in C_k} q_\eta(z)\, g(z) + q_\eta(\bar{C}_k)\, g(v), \qquad v \sim q_\eta \mid \bar{C}_k. \end{equation}

This is Rao-Blackwellization: unbiased, and variance never higher. And the one importance weight that survives is the scalar qη(Cˉk)q_\eta(\bar{C}_k) — constant in vv and bounded by 1, so it cannot explode the way a likelihood ratio can. Since LLM distributions concentrate most of their mass on a few tokens, qη(Cˉk)q_\eta(\bar{C}_k) is small and the estimator is nearly deterministic. Identifying CkC_k is free: it’s a top-kk over logits you already computed.

The caveat: this applies to terms you can evaluate at unsampled actions — KL and entropy regularizers, distillation against teacher logits, advantages from a value or Q model. For a pure outcome reward, Q(s,a)Q(s,a) at a different token requires rolling out a new continuation. Which is precisely what a tree does.

Rao-Blackwellizing the tree

Branch during sampling and the buffer becomes a tree: nodes are states, edges are tokens, leaves carry outcome rewards. GRPO’s group of NN completions is the trivial tree — branch NN ways at the root, never again. MCTS is the non-trivial one.

A tree gives extra samples of each of the four expectations in equation (6), but each must be Rao-Blackwellized with respect to the right part of the tree:

  1. Q(s,a)Q(s,a): average the subtree, not your own path. Every leaf below the edge (s,a)(s,a) is a sample of the future. Replacing a rollout’s own return with the mean return of its subtree is a conditional expectation given the tree — unbiased, variance never higher. This is exactly the MCTS backup.
  2. V(s)V(s): sibling subtrees, leave-one-out, at every node. The baseline applied to the edge (s,a)(s,a) must not depend on the data below that edge, or the zero-mean argument (equation 5) breaks and the estimator is biased. Use the leave-one-out mean of the sibling subtrees’ QQ estimates. GRPO’s group-mean baseline is exactly this at the root; in a tree you do it at every internal node, which gives each token a position-dependent baseline instead of one scalar per sequence.
  3. π(a∣s)\pi(a\mid s): children are action samples — or skip sampling. Multiple children of a node form a multi-sample Monte Carlo estimate of the action expectation, and branching is what turns the previous section’s caveat around: kk children means QQ evaluated at kk actions. The head-and-tail estimator (equation 8) composes with this — enumerate the high-probability children exactly, correct with the conditional tail.
  4. p(s)p(s): weight states by visitation, and don’t lose it when flattening. The marginal of a node is estimated by the fraction of rollouts passing through it — the number of leaves below it over NN, which is the MCTS visit count. When the tree is flattened into a packed batch, a shared prefix token appears once but stands for many rollouts, so its loss weight must carry that multiplicity. Deduplicate without reweighting and you have silently tilted p(s)p(s) toward deep, rare states — the tree version of GRPO’s per-sequence averaging mistake.

So the title isn’t a metaphor. The MCTS tree is the replay buffer — nodes are states, visit counts are the state marginal, backed-up values are the QQ estimates. And the buffer is the gradient estimator — flattened, it’s tokens, a mask, and per-token weights, where Rao-Blackwellization only changes the numbers in the weights. The training batch API doesn’t change at all, and the tree is the same prefix-sharing structure inference engines already exploit: compute the shared prefix once when sampling, once when training, and let the weights carry the multiplicity.

Experiment proposal

The claims above are falsifiable, so here is how to falsify them. Two tiers: an exact small-scale check where bias and variance can be measured against ground truth, and a real LLM run showing it translates to sample efficiency. The first makes people believe it; the second makes people adopt it.

Tier 1: ground-truth gradient comparison. Use a setting where the true policy gradient is exactly computable: a small transformer, vocabulary of 10–50 tokens, horizon short enough to enumerate the trajectory space (or near-enumerate with a huge-sample Monte Carlo reference), and a programmatic reward. At a few fixed checkpoints, draw many independent tree-buffers and compute each estimator’s gradient. Plot bias (error against the true gradient) and variance for:

  1. the single-rollout GRPO-style estimator,
  2. the tree with all four corrections,
  3. the tree with each correction knocked out one at a time: own-path return instead of subtree-mean QQ, a baseline that includes its own subtree instead of leave-one-out VV, deduplicated prefixes without visitation weights, and per-sequence instead of per-token averaging.

The math makes a sharp prediction. The full estimator is unbiased with strictly lower variance. The non-LOO baseline and the unweighted dedup are biased: their error should not shrink as samples are added. A figure where the knocked-out variants converge to the wrong gradient while the corrected estimator converges to the right one with smaller error bars is the whole argument in one plot. This tier is a day of compute on one GPU.

Tier 2: matched-compute LLM run. A small open model on math with verifiable reward — standard enough that the setup needs no defense. The key design decision is the compute accounting: compare at matched sampler FLOPs, not matched leaf count, because prefix sharing is half the point — a tree yields more leaves per FLOP. Three arms:

  1. GRPO with NN independent rollouts per prompt,
  2. tree rollouts (branch at a few entropy-based or random positions) with all four corrections,
  3. the same tree data processed naively: root-only baseline, own-path returns, no visitation weights.

Arm 3 is the important control — it separates “trees help” from “correct estimation on trees helps,” which is the actual claim. Report reward against sampler FLOPs, plus two mechanistic metrics during training: empirical advantage variance and gradient signal-to-noise across microbatches. If the math above is right, arm 2 dominates arm 1 in sample efficiency, and arm 3 underperforms or destabilizes late as its bias compounds.

A side experiment for the action expectation. Compare the exact per-token KL against a sampled estimator at equal compute and plot the variance of the KL term; same comparison for a head-and-tail distillation loss against sampled-token distillation.