中文博客 / 研究笔记
RL 中,训练器算的是 rollout 的那个 policy 吗?
从逐 token 的高层反馈,到 LastKV 的固定 chunk recurrence。
考虑一种很实际的递归语言模型 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 |
|---|---|---|
| 1 | 1 | 1 |
| 2 | 1.5 | 1.5 |
| 3 | 1.75 | 1.5 |
| 4 | 1.875 | 1.5 |
这只是说明依赖关系的算例,不是任何论文的模型实验。有限轮 refinement 只传播有限次反馈,串行计算则沿已生成的全部历史持续传播。若 logits 使用这些状态,条件概率就可能不同。足够多轮可以恢复这个严格因果递推,但那已经改变了固定少量并行 passes 的计算预算。
因此,真正的比较是:读取前一 token 已完成的递归状态,与读取某个较早 refinement 轮次的状态,是否是同一条依赖链。
LRT、FBT、T²MLR 如何讨论这个问题
这些论文已经研究了相关代价,不能简单地说它们没有意识到 mismatch,也不能把某个本地 adaptation 的规则当成所有原方法的规则。
- Latent Recurrent Transformer(LRT) 在 §4.6 直接指出:rollout 的递归状态以自回归方式生成,而策略概率由并行 refinement 重算,这会引入 mismatch。论文用缓存的 rollout 状态初始化重算时的状态 buffer,以减小差异;它宣称的是缓解,而非普遍的精确等价。
- Full-bandwidth Transformer(FBT) 在 §5 讨论 latent-feedback 轨迹的 replay,区分串行重算、hidden-state replay 与多轮 refinement。它把有限轮 replay 描述为对完整自回归分布的有限视野近似,也指出复用 behavior policy 的 hidden state 可能产生偏差。§5.2 的 RL 方案明确标注为尚未经过实验验证。
- T²MLR 在附录 B.2 讨论缓存 rollout 的精确 latent 与构建 BPTT 图。尤其需要注意,附录 B.5 说明主实验使用 exact sequential prompt prefill;Jacobi prefill 是以近似误差换取延迟的可选方案。
所以本文的前提是开头指定的执行组合:并行 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 结束处提前提交 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 的比较口径,逐步比较分布。
- 先查 logits 与完整分布。 报告逐位置的 log-prob 差异或 KL;相同 argmax 不足以证明一致。若 logits 只差一个公共常数,softmax 概率仍可以完全相同。
- 再查未更新时的 ratio。 使用真实 rollout old log-prob 作分母;不要让 replay 重算的分母掩盖执行差异。
- 专门覆盖边界。 对 \(C-1\)、\(C\)、\(C+1\) 和非整数 chunk 的 prompt,检查 prompt/response 切换、连续多次 chunk 提交及 t−1 blending。
- 独立检查梯度。 用小规模完整 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 内因果计算的并行性。
参考资料
- Latent Recurrent Transformer: Architecture Exploration, Training Strategies, and Scaling Behavior. arXiv:2605.26797v2。§4.4:prefill 与递归 decode;§4.6:RL mismatch 与 rollout state cache。
- Full-bandwidth Transformer. arXiv:2608.08888v2。§3.2:plain prompt 与 fused response;§5:latent-feedback replay、近似与递归梯度。
- T²MLR: Transformer with Temporal Middle-Layer Recurrence. arXiv:2607.15178v2。§2.4:并行近似;附录 B.2–B.5:RL、Jacobi 因果收敛与 prefill。
- Score Centering Stabilizes Off-policy Reinforcement Learning. arXiv:2609.20807v1。固定采样分布下的 score 均值与 drift 分解。
- LastKV: Last Layer Key-Value Cache is All You Need. 本站研究介绍。本文的 chunk 与 previous-chunk blending 规则按所讨论的 LastKV 设计定义。
- 模型为什么越走越偏:从 LLM 强化学习到自回归视频的 Drift。关于状态分布、梯度偏置与优化目标的背景讨论。