Blog / Research notes
In RL, Does the Trainer Evaluate the Policy That Generated the Rollout?
From higher-layer feedback at every token to LastKV's fixed chunk recurrence.
Consider a practical RL implementation for a recurrent language model. To control latency, the prompt uses parallel prefill. During response generation, each token reads the higher-layer state produced by the completed computation of the previous token, recurring one token at a time. During actor updates, the trainer preserves throughput by feeding the entire recorded sequence through a finite number of fully parallel refinement passes to recompute log-probabilities.
These three paths may not execute the same policy. Identical token sequences and parameters do not imply identical feedback states. The rollout uses states produced by the actual recurrence; the trainer uses states produced by a finite parallel approximation. They may assign different next-token probabilities to the same prefix.
In the previous article on drift in RL and autoregressive video, we distinguished conditional distribution mismatch, gradient bias, and optimization objectives. This article focuses on one specific issue in RL: whether sampling and gradient recomputation follow the same state update rules, and why LastKV fixes cross-layer feedback at chunk boundaries.
The central claim: finite parallel replay creates a structural policy mismatch if it cannot reproduce the feedback states of token-by-token rollout. LastKV defines chunk recurrence as part of the model’s behavior, allowing parallel recomputation within a chunk to correspond exactly to token-by-token decoding. The cost is that computation across chunks remains sequential.
Defining the consistency RL needs
Fix a prompt \(x\) and write the recorded response as \(a_{1:T}\). Let \(\mu_\theta\) denote the policy defined by the actual rollout path, and \(p_\theta\) the policy recomputed by the actor replay path. At identical parameters, we want each step to satisfy:
\[p_\theta(a_t\mid x,a_{<t}) =\mu_\theta(a_t\mid x,a_{<t}).\]A stronger check compares the entire next-token distribution, rather than only the probability of the token that happened to be sampled. This does not require updated parameters to remain equal to the old parameters. It requires both execution paths to express the same policy for any given set of parameters.
An ordinary causal Transformer provides a familiar example. Once the full trajectory is known, all positions can be computed in parallel under a causal mask. Token-by-token decoding instead caches historical KV at each layer. Under ideal numerical conditions, these are two execution orders of the same causal computation graph. Parallel training and sequential generation are therefore not inherently a problem.
The rollout considered here is a hybrid policy: the prompt passes through a specified parallel prefill, then the response continues with sequential recurrence from the state left by that prefill. Exact recomputation must reproduce the prefill, the state at the transition, and the subsequent recurrence. Replacing the whole sequence with a different parallel forward changes response states. Recomputing the whole sequence sequentially from the start may also change the initial state left by the prompt.
The issue is feedback state, not teacher forcing
Teacher forcing simply feeds known tokens to the model. In RL replay, those tokens come from a trajectory the model has already generated, rather than a separate labeled answer. Knowing the tokens does not mean knowing the recurrent states that the current parameters should produce. LastKV replay can also use teacher forcing.
A simple recurrence illustrates the distinction. Suppose the sequential model computes:
\[h_t=x_t+\rho h_{t-1},\qquad h_0=0.\]Parallel computation initializes \(h_t^{(0)}=x_t\) and performs the following update in each Jacobi refinement round:
\[h_t^{(k)}=x_t+\rho h_{t-1}^{(k-1)}, \qquad h_0^{(k)}=0.\]With \(x_t=1\) and \(\rho=0.5\), the difference between one refinement round and the actual sequential states is already visible:
| Position | Actual sequential state | One parallel refinement round after initialization |
|---|---|---|
| 1 | 1 | 1 |
| 2 | 1.5 | 1.5 |
| 3 | 1.75 | 1.5 |
| 4 | 1.875 | 1.5 |
This is an example of the dependency structure, not a model experiment from any paper. Finite refinement propagates feedback only a finite number of times, whereas sequential computation continues propagating it through the entire generated history. If logits depend on these states, the conditional probabilities may differ. Enough rounds can recover this strictly causal recurrence, but that changes the compute budget of a fixed small number of parallel passes.
The relevant comparison is therefore whether reading the previous token’s completed recurrent state and reading a state from an earlier refinement round follow the same dependency chain.
How LRT, FBT, and T²MLR discuss this issue
These papers already study the relevant tradeoffs. It would be inaccurate to say that they overlook mismatch, or to treat the rules of a local adaptation as the rules of every original method.
- Latent Recurrent Transformer (LRT) states directly in §4.6 that rollout recurrent states are generated autoregressively, while policy probabilities are recomputed through parallel refinement, introducing mismatch. The paper initializes the state buffer for recomputation with cached rollout states to reduce the discrepancy. Its claim is mitigation, rather than general exact equivalence.
- Full-bandwidth Transformer (FBT) discusses replay of latent-feedback trajectories in §5, distinguishing sequential recomputation, hidden-state replay, and multi-round refinement. It describes finite-round replay as a finite-horizon approximation to the full autoregressive distribution and notes that reusing behavior-policy hidden states may introduce bias. The RL proposal in §5.2 is explicitly marked as not yet experimentally validated.
- T²MLR discusses caching exact rollout latents and constructing a BPTT graph in Appendix B.2. In particular, Appendix B.5 states that the main experiments use exact sequential prompt prefill. Jacobi prefill is an optional method that trades approximation error for latency.
The premise of this article is thus the execution combination specified at the start: parallel prefill, token-by-token recurrent generation, and a finite number of fully parallel actor replay passes. It is not a literature claim that all three papers always use parallel prefill. An approximation may work well on measured tasks, but similar average losses or benchmark scores do not establish equal probabilities at each prefix.
Why it affects policy gradients
To optimize the expected reward of the actual rollout policy, the objective and its score-function gradient are:
\[\begin{aligned} J(\theta)&=\mathbb E_{\tau\sim\mu_\theta}[R(\tau)],\\ \nabla_\theta J &=\mathbb E_{\tau\sim\mu_\theta}\!\left[\right.\\ &\qquad\left.R(\tau)\nabla_\theta\log\mu_\theta(\tau)\right]. \end{aligned}\]Here we omit environment terms that do not depend on the parameters and assume the reward has no additional explicit parameter dependence. If the update instead uses \(\nabla_\theta\log p_\theta(\tau)\), it cannot generally be identified with this gradient. The recurrent policy generates the data, while differentiation passes through the parallel approximation policy.
This discrepancy can occur even before the first optimizer update. If a PPO-style probability ratio stores the actual rollout old log-probability in its denominator and uses actor replay for its numerator, then at \(\theta=\theta_{\rm old}\):
\[r_t= \frac{p_{\theta_{\rm old}}(a_t\mid x,a_{<t})} {\mu_{\theta_{\rm old}}(a_t\mid x,a_{<t})}\]may already differ from 1. This can still be a valid density ratio between two distributions, but it includes a change of execution mode, rather than the intended ratio before and after updating the same deployed policy. Recomputing the denominator with the parallel engine as well can make the initial ratio equal to 1 while concealing the actual sampling discrepancy.
This connects to the previous discussion of Score Centering. At a fixed prefix, let \(g=\nabla_\theta\log p_\theta\). An uncentered score update decomposes as:
\[\mathbb E_\mu[Rg] =\mathbb E_\mu[R]\,\mathbb E_\mu[g] +\operatorname{Cov}_\mu(R,g).\]When the sampler and trainer differ, \(\mathbb E_\mu[g]\) is generally no longer zero. Centering can subtract this mean-score term, but it does not automatically turn the remaining term into the exact policy gradient of \(\mu_\theta\). This explains one source of bias; it does not imply that all RL implementations using baselines, clipping, or within-group normalization will drift in the same way.
LastKV: making recurrence boundaries part of the policy
LastKV chooses a different set of state update rules. It does not require every token’s last-layer state to feed back immediately to all lower layers of the next token. It defers this cross-layer feedback to fixed chunk boundaries.
Start with the core rule without blending. Within chunk \(c\), layer \(\ell\) reads last-layer KV from completed chunks as history and uses its own per-layer KV for the current chunk:
\[A_c^\ell= \operatorname{Attn}\left( Q_c^\ell, [K_{<c}^{L};K_c^\ell], [V_{<c}^{L};V_c^\ell] \right).\]All history precedes the current chunk, and a causal mask applies to the current portion. The rules for reading cross-chunk history remain fixed while this chunk is processed. A token within the current chunk is not promoted to historical feedback that every layer immediately reads merely because it has completed the last layer.
Prefill and actor replay therefore process chunks sequentially, with parallel computation inside each chunk. Decoding samples one token at a time, but grows the local KV of each layer within the chunk. For identical input tokens, causal attention ensures that earlier positions are unaffected by later positions that have not yet been generated. The chunk’s role in subsequent history changes only once the chunk is complete.
Actor replay
Fixed history → parallel causal forward within the chunk → complete chunk → update history → next chunk
Rollout decode
Same history → generate tokens and grow per-layer KV within the chunk → update history at the same boundary → next chunk
Both paths can execute the same causal graph. Exactness here means mathematical forward equivalence for given parameters and a given trajectory. Different kernels, precision, dropout, and cache implementations still require separate checks.
With previous-chunk blending, the immediately preceding completed chunk also retains its per-layer KV, mixing it with last-layer KV through gates per layer and per KV head:
\[\widetilde K_{c-1}^\ell =(1-\alpha_\ell)K_{c-1}^\ell+\alpha_\ell K_{c-1}^{L}, \qquad \widetilde V_{c-1}^\ell =(1-\alpha_\ell)V_{c-1}^\ell+\alpha_\ell V_{c-1}^{L}.\]A newly completed chunk therefore does not immediately discard all per-layer KV. It participates in the blend for the next chunk; only when it becomes older history does it use last-layer KV alone. This remains part of the same state transition as long as rollout and replay use identical blending and transition timing.
The end of a prompt is not a chunk boundary
Suppose the chunk size is \(C=512\) and the prompt contains 600 tokens. The first chunk is complete, while the second contains only 88 tokens. Response generation should continue from this partial chunk’s per-layer KV, positions, and offset. The second chunk completes after 424 more tokens are processed, when the total input length reaches 1024.
The count here is the number of tokens fed into and fully processed by the model. The last prompt token’s logits can already be used to sample the first response token. Sampling alone has not yet written that response token into the KV cache.
- Fully filled chunks
- 1
- Tokens in current chunk
- 88 / 512
- Tokens to next boundary
- 424
- Total tokens at next boundary
- 1024
The most recent full chunk keeps its per-layer KV for t−1 blending in the next chunk; it does not immediately become pure last-layer KV.
If an inference service commits the partial chunk early at the end of the prompt, or the trainer starts the response in a new chunk, it changes when feedback becomes visible. This behavior could be defined as another explicit policy, but replay would have to reproduce it. It cannot be treated as a free implementation adjustment under fixed chunk rules.
Likewise, padding, packed documents, truncation, and minibatch slicing must not silently reset the chunk phase. Matching the chunk size is only the first step. Absolute boundaries, document resets, RoPE positions, and cache transition rules must also match.
Matching probabilities does not guarantee complete gradients
Caching rollout latents and reusing them during training may bring states closer. But two questions must be answered separately: whether the current forward uses the correct states, and whether gradients pass through the computation that produced them.
If \(h_t=F_\theta(h_{t-1},a_t)\), the log-probability derivative generally includes both a direct parameter term and a historical state term:
\[\begin{aligned} \frac{d\log\pi_\theta}{d\theta} &=\frac{\partial\log\pi_\theta}{\partial\theta}\\ &\quad+\frac{\partial\log\pi_\theta}{\partial h} \frac{dh}{d\theta}. \end{aligned}\]Before parameters are updated, reusing exact rollout states may align the forward computation. Detaching those states as constants, however, omits the corresponding recurrent derivatives. Continuing to reuse old latents or KV after a parameter update may also introduce forward discrepancies from stale states.
LastKV is subject to the same distinction. Detaching the historical KV of completed chunks can leave logits unchanged while removing gradient paths from future losses through the producers of historical KV. Complete recurrent gradients require preserving or recomputing the corresponding computation graph; TBPTT is another approximation that can be chosen explicitly. Forward policy mismatch and backward gradient truncation are separate problems. A single claim of consistency cannot cover both.
How to verify consistency, and what this design costs
The most direct check does not require a full RL experiment. Fix a checkpoint, record a rollout, and recompute the same trajectory with the actor’s training forward. Disable dropout, align masks, positions, and precision, and use the same comparison convention for temperature and logits processing. Then compare distributions step by step.
- Check logits and full distributions first. Report per-position log-probability differences or KL. Matching argmax is insufficient to establish consistency. If logits differ only by a common additive constant, softmax probabilities can still be identical.
- Check the ratio before any update. Use the actual rollout old log-probabilities in the denominator. A denominator recomputed by replay must not conceal execution differences.
- Cover boundaries explicitly. For prompts of length \(C-1\), \(C\), \(C+1\), and nonintegral numbers of chunks, check the prompt/response transition, consecutive chunk commits, and t−1 blending.
- Check gradients independently. Use a small complete autograd reference to verify paths across chunks. If gradients are truncated, report the truncation scope explicitly.
LastKV’s benefits and costs can be compared in the same table:
| Execution rule | Actor replay | Relationship to rollout |
|---|---|---|
| Ordinary causal Transformer | Parallel causal forward over the full sequence | Can correspond exactly to token-by-token cached decoding |
| Higher-layer feedback at every token + finite parallel approximation | Multi-round refinement over the full sequence | Generally approximates actual sequential feedback; state differences need checking |
| Exact reproduction of the hybrid recurrent trajectory | Match prefill, then recompute response states sequentially | Can preserve forward consistency while retaining temporal dependencies |
| LastKV with fixed chunk recurrence | Sequential across chunks, parallel within each chunk | Can preserve forward consistency under the same chunk rules |
LastKV does not obtain full-sequence parallelism and token-by-token deep feedback for free. An input of length \(N\) still has approximately \(\lceil N/C\rceil\) sequential chunk stages. Larger chunks offer more parallelism while deferring higher-layer feedback longer. Token-by-token sampling itself remains sequential, and inference KV storage does not directly determine training memory when the recurrent graph is retained.
This article analyzes architecture and execution rules; it reports no new LastKV RL experiments or serving benchmarks. Actual implementations still need to verify cached decoding, replay, and backward paths. The design advantage is to first define recurrence at chunk granularity, share those rules between rollout and gradient recomputation, and then exploit parallel causal computation within each chunk.
References
- Latent Recurrent Transformer: Architecture Exploration, Training Strategies, and Scaling Behavior. arXiv:2605.26797v2. §4.4: prefill and recurrent decoding; §4.6: RL mismatch and rollout state caching.
- Full-bandwidth Transformer. arXiv:2608.08888v2. §3.2: plain prompt and fused response; §5: latent-feedback replay, approximations, and recurrent gradients.
- T²MLR: Transformer with Temporal Middle-Layer Recurrence. arXiv:2607.15178v2. §2.4: parallel approximation; Appendices B.2–B.5: RL, Jacobi causal convergence, and prefill.
- Score Centering Stabilizes Off-policy Reinforcement Learning. arXiv:2609.20807v1. Mean score and drift decomposition under a fixed sampling distribution.
- LastKV: Last Layer Key-Value Cache is All You Need. Research overview on this site. Chunk and previous-chunk blending rules here follow the LastKV design under discussion.
- Why Models Drift: From LLM Reinforcement Learning to Autoregressive Video. Background on state distributions, gradient bias, and optimization objectives.