中文博客 / 研究笔记

RL 中,训练器算的是 rollout 的那个 policy 吗?

从逐 token 的高层反馈,到 LastKV 的固定 chunk recurrence。

Zefan Cai · · 中文研究笔记

考虑一种很实际的递归语言模型 RL 实现:为了控制延迟,prompt 使用并行 prefill;生成 response 时,当前 token 读取前一个 token 完成计算后的高层状态,逐 token 递归;更新 actor 时,为了保留训练吞吐,又把记录下来的整条序列送进有限轮的全并行 refinement,重算 log-prob。

这三条路径可能没有在执行同一个 policy。 token 序列和参数相同,不代表反馈状态也相同。rollout 用真实递归得到的状态,训练器用有限轮并行近似得到的状态,二者就可能给同一个 prefix 分配不同的下一 token 概率。

在上一篇关于 RL 与自回归视频 drift 的文章里,我们区分了条件分布偏移、梯度偏置和优化目标。这篇只讨论 RL 中更具体的一处:采样与梯度重算是否遵循同一套状态更新规则,以及 LastKV 为什么把跨层反馈固定在 chunk 边界。

核心判断:有限轮并行 replay 若不能复现逐 token rollout 的反馈状态,就存在结构性的 policy mismatch。LastKV 把 chunk recurrence 定义为模型本身的行为,使 chunk 内并行重算与逐 token decode 可以精确对应;代价是跨 chunk 仍需顺序执行。

先定义 RL 需要的一致性

固定 prompt \(x\),把记录下来的 response 写成 \(a_{1:T}\)。记 \(\mu_\theta\) 为实际 rollout 路径定义的策略,\(p_\theta\) 为 actor replay 路径重算出的策略。在相同参数下,我们希望每一步满足:

\[p_\theta(a_t\mid x,a_{<t}) =\mu_\theta(a_t\mid x,a_{<t}).\]

更强的检查是整个下一 token 分布相同,而不仅是碰巧采样到的那个 token 概率相同。这也不要求更新后的参数仍等于旧参数;它要求的是,对于任意给定的一组参数,两条执行路径表达同一个策略。

普通 causal Transformer 就是熟悉的例子:已知整条轨迹后,可以在 causal mask 下并行计算所有位置;逐 token decode 则缓存各层历史 KV。理想数值条件下,它们是同一个因果计算图的两种执行顺序。因此,“并行训练、串行生成”本身并不构成问题。

这里的 rollout 是一个 hybrid policy:prompt 经过指定的并行 prefill,response 从该 prefill 留下的状态继续串行递归。要精确重算,就必须复现这段 prefill、切换时的状态,以及后续递归。把整段改成另一个并行 forward 会改变 response 状态;把整段从头改成串行计算,也可能改变 prompt 留下的初始状态。

问题来自反馈状态,不是 teacher forcing

teacher forcing 只是把已知 token 喂给模型。RL replay 中这些 token 来自模型已经生成的轨迹,而不是另一份标注答案。知道 token,并不等于知道当前参数下应该产生的递归状态。 LastKV 的 replay 同样可以使用 teacher forcing。

用一个简化递推说明区别。设串行模型执行:

\[h_t=x_t+\rho h_{t-1},\qquad h_0=0.\]

并行初始化为 \(h_t^{(0)}=x_t\),每轮 Jacobi refinement 执行:

\[h_t^{(k)}=x_t+\rho h_{t-1}^{(k-1)}, \qquad h_0^{(k)}=0.\]

取 \(x_t=1\)、\(\rho=0.5\),一轮 refinement 与真实串行状态的区别已经可见:

位置真实串行状态初始化后的一轮并行 refinement
111
21.51.5
31.751.5
41.8751.5

这只是说明依赖关系的算例,不是任何论文的模型实验。有限轮 refinement 只传播有限次反馈,串行计算则沿已生成的全部历史持续传播。若 logits 使用这些状态,条件概率就可能不同。足够多轮可以恢复这个严格因果递推,但那已经改变了固定少量并行 passes 的计算预算。

因此,真正的比较是:读取前一 token 已完成的递归状态,与读取某个较早 refinement 轮次的状态,是否是同一条依赖链。

LRT、FBT、T²MLR 如何讨论这个问题

这些论文已经研究了相关代价,不能简单地说它们没有意识到 mismatch,也不能把某个本地 adaptation 的规则当成所有原方法的规则。

所以本文的前提是开头指定的执行组合:并行 prefill、逐 token 递归生成、有限轮全并行 actor replay。 它不是“三篇论文都固定使用并行 prefill”的文献结论。近似可能在实测任务上很好,但相近的平均 loss 或 benchmark 分数不证明逐 prefix 的概率相同。

为什么它影响 policy gradient

如果想优化实际 rollout 策略的期望奖励,目标与相应的 score-function 梯度是:

\[\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}\]

这里省略与参数无关的环境项,并假设奖励没有额外的显式参数依赖。若实际更新改用 \(\nabla_\theta\log p_\theta(\tau)\),一般不能把它等同于上面的梯度。数据由递归策略生成,求导却经过并行近似策略。

这个差异甚至可以发生在第一次 optimizer update 之前。若 PPO 风格的概率比分母保存了真正的 rollout old log-prob,分子用 actor replay,那么在 \(\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})}\]

也可能偏离 1。这个比值仍然可以是两个分布的合法密度比,但它包含了执行模式的切换,不是预期的“同一部署策略更新前后”的比值。若分母也由并行引擎重算,初始 ratio 等于 1 又可能把真正的采样差异藏起来。

这与前文的 Score Centering 讨论相接。在固定 prefix 下,令 \(g=\nabla_\theta\log p_\theta\),未经中心化的 score 更新可分解为:

\[\mathbb E_\mu[Rg] =\mathbb E_\mu[R]\,\mathbb E_\mu[g] +\operatorname{Cov}_\mu(R,g).\]

当 sampler 与 trainer 不一致时,\(\mathbb E_\mu[g]\) 通常不再为零。中心化可以扣掉这里的平均 score 项,但不自动把剩余项变成 \(\mu_\theta\) 的精确策略梯度。这解释一种偏置来源,并不意味着所有带 baseline、clipping 或组内归一化的 RL 实现都会以同样方式漂移。

LastKV:把递归边界写进 policy

LastKV 选择另一套状态更新规则。它不要求每一个 token 的最后层状态立即反馈给下一个 token 的所有低层,而是把这次跨层反馈推迟到固定的 chunk 边界。

先看不带 blending 的核心规则。对 chunk \(c\) 中的第 \(\ell\) 层,历史由完整 chunk 的最后层 KV 构成,当前 chunk 则使用本层自己的 KV:

\[A_c^\ell= \operatorname{Attn}\left( Q_c^\ell, [K_{<c}^{L};K_c^\ell], [V_{<c}^{L};V_c^\ell] \right).\]

历史完整地位于当前 chunk 之前,当前部分施加 causal mask。在处理这个 chunk 的过程中,跨 chunk 的历史读取规则保持固定。 当前 chunk 内也不会因为某个 token 已经过完最后层,就把它升级成所有层立即读取的历史反馈。

这样,prefill 与 actor replay 都按 chunk 顺序执行,每个 chunk 内并行计算;decode 虽然逐 token 采样,但在 chunk 内增长的是各层自己的局部 KV。对相同的输入 token,causal attention 保证早期位置不受尚未生成的后续位置影响。直到 chunk 完成,才更新它在后续计算中的历史角色。

Actor replay

固定历史 → chunk 内并行 causal forward → 完成 chunk → 更新历史 → 下一 chunk

Rollout decode

同一历史 → chunk 内逐 token 生成并增长本层 KV → 同一边界更新历史 → 下一 chunk

两条路径可以执行同一个因果图。这里的精确性指给定参数和轨迹时的数学前向等价;不同 kernel、精度、dropout 或 cache 实现仍需单独核验。

实际使用 previous-chunk blending 时,紧邻的完整 chunk 还保留本层 KV,并按每层、每 KV head 的门控与最后层 KV 混合:

\[\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}.\]

所以刚完成的 chunk 不会立即丢掉所有本层 KV:它在下一 chunk 中参与 blend;再成为更老的历史时,才只使用最后层 KV。只要 rollout 与 replay 保持相同的 blending 和切换时机,这也属于同一套状态转移。

Prompt 结束不是 chunk 边界

设 chunk size 为 \(C=512\),prompt 长度为 600。第一个 chunk 已完成,第二个 chunk 只填了 88 个 token。生成 response 时,应继续使用这个 partial chunk 的本层 KV、位置与 offset;再处理 424 个 token,总输入长度达到 1024,才完成第二个 chunk。

这里按已送入模型并完成处理的 token 数计数。最后一个 prompt token 的 logits 已可用于采样第一个 response token;采样动作本身还没有把该 response token 写进 KV。

已填满的 chunk
1
当前 chunk 已填入
88 / 512
距下一边界还需
424
下一边界的总 token 数
1024

刚完成的 chunk 在下一 chunk 中仍保留本层 KV 参与 t−1 blending;它不会立即变成纯最后层 KV。

Prompt后续 response(示意)
Chunk 1 · 1–512
Prompt 512 + 后续 response 0
Chunk 2 · 513–1024 · 当前
Prompt 88 + 后续 response 424
Chunk 3 · 1025–1536
Prompt 0 + 后续 response 512
Chunk 4 · 1537–2048
Prompt 0 + 后续 response 512
固定 C = 512,展示前 2048 个 token 的位置划分。response 颜色只表示后续位置,不是模型生成结果。拖动 prompt 长度,chunk 边界仍固定在 512 的倍数;关闭 JavaScript 时保留默认示例。

如果推理服务在 prompt 结束处提前提交 partial chunk,或者训练器把 response 从一个新 chunk 开始,便改变了反馈的可见时机。这样的行为可以被定义成另一套明确的 policy,但 replay 必须复现它;不能把它当成固定 chunk 规则下免费的实现调整。

同理,padding、packed documents、截断和 minibatch 切片也不能悄悄重设 chunk 的相位。匹配 chunk size 只是第一步;绝对边界、文档 reset、RoPE 位置与缓存切换规则也要匹配。

概率一致,还不等于梯度完整

缓存 rollout 的 latent,再在训练时复用,看起来可以让状态更接近。但需要分开回答两个问题:当前 forward 使用的状态是否正确,以及梯度是否经过产生这些状态的计算。

若 \(h_t=F_\theta(h_{t-1},a_t)\),log-prob 的导数一般包括直接参数项与历史状态项:

\[\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}\]

在尚未更新参数时,复用精确 rollout 状态可能让 forward 对上;把状态当常量 detach,却会省去相应的递归导数。更新参数后继续复用旧 latent 或 KV,还可能引入状态陈旧导致的前向差异。

LastKV 也不例外。detach 完整 chunk 的历史 KV 可以保持 logits 不变,同时删掉未来损失经过历史 KV producer 的梯度路径。要获得完整的递归梯度,需要保留或重算相应计算图;TBPTT 是另一个可以明确选择的近似。前向 policy mismatch 与反向梯度截断是两类问题,不能用同一个“已经一致”概括。

如何验证,以及这项设计的代价

最直接的检查不需要先跑完整 RL 实验:固定 checkpoint,记录一条 rollout,然后用 actor 的训练 forward 重算同一轨迹。关闭 dropout,统一 mask、位置、精度及 temperature / logits processing 的比较口径,逐步比较分布。

  1. 先查 logits 与完整分布。 报告逐位置的 log-prob 差异或 KL;相同 argmax 不足以证明一致。若 logits 只差一个公共常数,softmax 概率仍可以完全相同。
  2. 再查未更新时的 ratio。 使用真实 rollout old log-prob 作分母;不要让 replay 重算的分母掩盖执行差异。
  3. 专门覆盖边界。 对 \(C-1\)、\(C\)、\(C+1\) 和非整数 chunk 的 prompt,检查 prompt/response 切换、连续多次 chunk 提交及 t−1 blending。
  4. 独立检查梯度。 用小规模完整 autograd reference,验证跨 chunk 路径;若使用截断,就明确报告截断范围。

LastKV 的收益与代价可以放在同一张表里:

执行规则Actor replay与 rollout 的关系
普通 causal Transformer整段并行 causal forward可精确对应逐 token cached decode
逐 token 高层反馈 + 有限轮并行近似全序列、多轮 refinement一般近似真实串行反馈;需检查状态差异
精确复现 hybrid 递归轨迹匹配 prefill,再串行重算 response 状态可保持前向一致,但保留时间依赖
LastKV 固定 chunk recurrence跨 chunk 顺序、chunk 内并行同一 chunk 规则下可保持前向一致

LastKV 没有获得“整段全并行,又保留逐 token 深层反馈”的免费组合。 长度为 \(N\) 的输入仍有约 \(\lceil N/C\rceil\) 个顺序 chunk 阶段;更大的 chunk 提供更多并行性,也把高层反馈推迟得更久。逐 token 采样本身仍然串行,保留递归图的训练内存也不能由推理 KV 存储量直接推断。

本文给出的是架构与执行规则的分析,不报告新的 LastKV RL 实验或 serving benchmark。实际实现仍需验证 cached decode、replay 与反向路径。它的设计优势是:先把递归粒度定义成 chunk,让 rollout 与梯度重算共享这套规则,再利用 chunk 内因果计算的并行性。

参考资料

  1. Latent Recurrent Transformer: Architecture Exploration, Training Strategies, and Scaling Behavior. arXiv:2605.26797v2。§4.4:prefill 与递归 decode;§4.6:RL mismatch 与 rollout state cache。
  2. Full-bandwidth Transformer. arXiv:2608.08888v2。§3.2:plain prompt 与 fused response;§5:latent-feedback replay、近似与递归梯度。
  3. T²MLR: Transformer with Temporal Middle-Layer Recurrence. arXiv:2607.15178v2。§2.4:并行近似;附录 B.2–B.5:RL、Jacobi 因果收敛与 prefill。
  4. Score Centering Stabilizes Off-policy Reinforcement Learning. arXiv:2609.20807v1。固定采样分布下的 score 均值与 drift 分解。
  5. LastKV: Last Layer Key-Value Cache is All You Need. 本站研究介绍。本文的 chunk 与 previous-chunk blending 规则按所讨论的 LastKV 设计定义。
  6. 模型为什么越走越偏:从 LLM 强化学习到自回归视频的 Drift。关于状态分布、梯度偏置与优化目标的背景讨论。

← 所有博客 · 回到正文