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: