OPD 随想

引言

笔者最近在做 MOPD 研究时发现,在计算 Advantage(教师-学生预测概率的 log-ratio)时,仅用学生采样的 top-1 token 计算已经效果很好了。巧合的是,Kimi K3 技术报告中也称,top-k 采样的消融未见明显性能或效率提升。因此有感而发,在本文中讨论一个问题:在 On-Policy Distillation 中,为什么不应该多步在全词表上计算 Advantage?


背景:On-Policy Distillation 的目标

在线蒸馏的目标是让 student \(q_\theta\) 逼近 teacher \(p\),并且是在 student 的采样轨迹上,以防止暴露偏差(exposure bias)或分布不匹配。从信息论角度,最自然的目标是最小化 reverse KL 散度:

\[\mathcal{L}(\theta) = \mathbb{E}_{x \sim p_x}\left[\mathrm{KL}\!\left[q_\theta(\cdot|x) \,\|\, p(\cdot|x)\right]\right]\]

由于序列空间的维度灾难,直接计算该期望不可行,需通过采样估计梯度。对序列生成任务展开:

\[\mathcal{L}(\theta) = \mathbb{E}_{x \sim p_x,\; y \sim q_\theta}\left[\sum_{t=1}^{T} \log \frac{q_\theta(y_t|y_{<t}, x)}{p(y_t|y_{<t}, x)}\right]\]

对 \(\theta\) 求梯度时,注意到损失函数中有两处 \(\theta\) 的依赖:采样分布 \(y \sim q_\theta\),以及被期望的函数本身也含 \(\theta\)。因此梯度由两项构成:

\[\nabla_\theta \mathcal{L}(\theta) = \underbrace{\mathbb{E}_{y\sim q_\theta}\!\left[f(y;\theta)\,\nabla_\theta\log q_\theta(y)\right]}_{\text{log-trick 项}} + \underbrace{\mathbb{E}_{y\sim q_\theta}\!\left[\nabla_\theta f(y;\theta)\right]}_{\text{直接梯度项}}\]

其中直接梯度项为:

\[\mathbb{E}_{y\sim q_\theta}\!\left[\nabla_\theta f(y;\theta)\right] = \mathbb{E}_{y\sim q_\theta}\!\left[\sum_{t=1}^T \nabla_\theta \log q_\theta(y_t|y_{<t},x)\right]\]

合并两项,得到完整的 Policy Gradient 形式梯度:

\[\nabla_\theta \mathcal{L}(\theta) = \mathbb{E}_{y\sim q_\theta}\!\left[\sum_{t=1}^T (R_t - 1)\cdot \nabla_\theta \log q_\theta(y_t|y_{<t},x)\right]\]

其中即时奖励与多步 Return 分别定义为:

\[r_t = \log \frac{q_\theta(y_t|y_{<t}, x)}{p(y_t|y_{<t}, x)}, \qquad R_t = \sum_{t'=t}^{T} r_{t'}\]

这是标准的 REINFORCE 结构,\(R_t\) 为从第 \(t\) 步起的累积奖励,\(-1\) 项来自 \(f\) 对 \(\theta\) 的直接偏导,等价于一个常数 baseline,不影响梯度方向,但影响方差。


问题的提出

在实现中,一个看似自然的想法是:

既然每个 token 位置 \(t\) 上,\(q_\theta\) 已经计算了整个词表的概率分布,能否不采样单个 \(y_t\),而是对全词表 \(V\) 枚举,用加权求和替代采样,从而降低方差?

形式上,这将梯度估计从单点采样

\[\hat{g}_t^{\text{sample}} = R_t \cdot \nabla_\theta \log q_\theta(y_t|y_{<t}, x), \quad y_t \sim q_\theta\]

替换为全词表期望

\[\hat{g}_t^{\text{vocab}} = \sum_{v \in V} R_t(v) \cdot \nabla_\theta \log q_\theta(v|y_{<t}, x) \cdot q_\theta(v|y_{<t}, x)\]

这个替换看起来用期望取代了采样,方差更低,实则从根本上破坏了轨迹一致性


核心问题:轨迹一致性

什么是合法的轨迹?

Policy Gradient 的理论基础要求梯度估计量作用在从 \(q_\theta\) 真实采样的完整轨迹上。一条合法轨迹为:

\[y = (y_1, y_2, \ldots, y_T), \quad y_t \sim q_\theta(\cdot|y_{<t}, x)\]

每个 \(y_t\) 与后续 \(y_{t+1}, \ldots, y_T\) 之间存在真实的因果依赖关系:\(y_{t+1}\) 是在已知 \(y_t\) 的条件下采样的。

多步 Return 的语义

\(R_t\) 的含义是:在执行动作 \(y_t\) 之后,沿当前策略继续走完整条轨迹所获得的累积奖励

\[R_t = r_t(y_t) + r_{t+1}(y_{t+1}) + \cdots + r_T(y_T)\]

其中 \(y_{t+1}, \ldots, y_T\) 是在 \(y_t\) 确定之后继续从 \(q_\theta\) 采样得到的。因此 \(R_t\) 是 \((y_t, y_{t+1}, \ldots, y_T)\) 的联合函数,记作 \(R_t(y_t, y_{t+1}, \ldots, y_T)\)。

全词表枚举时发生了什么?

在全词表枚举中,对每个候选 \(v \in V\),需要给出 \(R_t(v)\)。但问题在于:\(y_{t+1}, \ldots, y_T\) 从哪里来?实践中只有两种方案,两种均有致命缺陷

方案 A:固定后续轨迹(有偏)

用同一条采样轨迹的后续部分 \(\hat{y}_{t+1}, \ldots, \hat{y}_T\) 对所有候选 \(v\) 共用:

\[R_t(v) \approx r_t(v) + \sum_{t'=t+1}^{T} r_{t'}(\hat{y}_{t'})\]
问题:\(\hat{y}_{t+1}\) 是在 \(\hat{y}_t\) 的条件下采样的,即 $$\hat{y}{t+1} \sim q\theta(\cdot \hat{y}t, y{<t}, x)\(。当候选换成\)v \neq \hat{y}_t\(时,这条后续轨迹对\)v$$ 而言并不合法。

数学上,正确的期望分解为:

\[\mathbb{E}_{y_t, y_{>t}}[R_t] = \mathbb{E}_{y_t}\!\left[\mathbb{E}_{y_{>t}|y_t}[R_t]\right]\]

而选择 A 实际计算的是:

\[\sum_v q_\theta(v) \cdot \left[r_t(v) + \underbrace{\sum_{t'>t} r_{t'}(\hat{y}_{t'})}_{\text{与 }v\text{ 无关的常数}}\right]\]

后半部分是与 \(v\) 无关的常数,对梯度的贡献退化;同时 \(\hat{y}_{t+1}\) 条件于错误的上文,产生系统性偏差

方案 B:对每个 \(v\) 独立 rollout(不可行)

对每个 \(v \in V\),从 \(v\) 出发重新采样 \(y_{t+1}^{(v)}, \ldots, y_T^{(v)}\),计算各自的 \(R_t^{(v)}\):

\[\hat{g}_t^{\text{vocab}} = \sum_{v \in V} R_t^{(v)} \cdot \nabla_\theta \log q_\theta(v|y_{<t}, x) \cdot q_\theta(v|y_{<t}, x)\]

这在数学上无偏,但代价是:

  • 词表大小 $$ V \sim 10^5\(,序列长度\)T \sim 10^3$$
  • 每个位置需要 $$ V $$ 次完整 rollout
  • 总计算量:$$O( V \cdot T)$$ 次完整前向传播

在任何实际系统中均不可行。


如何正确地进行全词表对齐?

值得特别指出,以下操作完全合法

\[\hat{g}_t^{\text{instant}} = \sum_{v \in V} r_t(v) \cdot \nabla_\theta \log q_\theta(v|y_{<t}, x) \cdot q_\theta(v|y_{<t}, x)\]
其中 $$r_t(v) = \log \frac{p(v y_{<t}, x)}{q_\theta(v y_{<t}, x)}$$ 仅是单步即时奖励,不涉及任何后续轨迹。

该量存在精确闭合形式。注意到:

\[\sum_{v \in V} \log \frac{p(v|y_{<t},x)}{q_\theta(v|y_{<t},x)} \cdot \nabla_\theta \log q_\theta(v|y_{<t},x) \cdot q_\theta(v|y_{<t},x) = -\nabla_\theta \,\mathrm{KL}\!\left[q_\theta(\cdot|y_{<t},x) \,\|\, p(\cdot|y_{<t},x)\right]\]

这正是 token-level KL 蒸馏(SeqKD)的直接梯度

它合法的根本原因在于:即时奖励 \(r_t(v)\) 完全由当前位置的概率决定,不依赖任何后续轨迹,因此全词表枚举不产生轨迹不一致性问题。


偏差-方差视角的完整对比

方法 偏差 方差 计算量 备注
采样 + 即时 reward 无偏 \(O(T)\) 高方差,需要 baseline
全词表 + 即时 reward 无偏 \(O(|V| \cdot T)\) 即 SeqKD,推荐
采样 + 多步 return 无偏 \(O(T)\) 标准 PG / REINFORCE
全词表 + 固定轨迹 return 有偏 \(O(|V| \cdot T)\) ❌ 错误做法
全词表 + 各自 rollout 无偏 \(O(|V|^T)\) 理论正确,实践不可行

正确的梯度结构

综合以上分析,On-Policy Distillation 中合法的梯度估计量有两种形式:

形式一:REINFORCE(RL-style 训练)

\[\hat{g} = \sum_{t=1}^{T} R_t \cdot \nabla_\theta \log q_\theta(y_t|y_{<t}, x), \quad y \sim q_\theta\]

其中 \(y_t\) 是真实采样的 token,\(R_t\) 建立在与 \(y_t\) 具有合法因果关系的后续轨迹之上。

形式二:Token-level KL(GKD 范式)

\[\hat{g} = \sum_{t=1}^{T} \nabla_\theta \,\mathrm{KL}\!\left[q_\theta(\cdot|y_{<t},x) \,\|\, p(\cdot|y_{<t},x)\right]\]

其中 \(y_{<t}\) 可来自 \(q_\theta\) 的 rollout(on-policy)或固定数据集(off-policy)。


总结

不应该对全词表计算多步 Advantage,原因可以精确表述为:

多步 Return \(R_t\) 是当前 token \(y_t\) 与后续轨迹 \(y_{>t}\) 的联合函数。全词表枚举时,若对每个候选 \(v\) 共用同一条后续轨迹,则该后续轨迹与 \(v\) 之间不存在合法的因果关系,导致 return 的计算语义被破坏,产生有偏的梯度估计;若为每个 \(v\) 单独 rollout,则计算量根本不可行。

这一约束与”改变随机变量的身份”无关——全词表求和在数学上就是期望操作,本身完全合法。真正的约束来自轨迹一致性:多步 return 要求 token 与其后续轨迹之间存在真实的条件依赖关系,而




Enjoy Reading This Article?

Here are some more articles you might like to read next:

  • LU-KV: KV Cache Optimization Based on Long-term Utility