docs: 修正 π̂ detach 失败机制的方向性错误(压低先验逃逸而非抬高自我强化),补充推导要点
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
@@ -58,7 +58,7 @@ conda activate ars-opd && ruff check ars_opd/ --fix && ruff format ars_opd/
|
||||
| 类型 | 要求 | 示例 |
|
||||
|------|------|------|
|
||||
| 论文锚点 | 实现论文公式/机制的函数,docstring 首行标出处;关键行旁给公式本体 | `# 式(5): π̂ = (k_sem + α·π̄) / (N + α)` |
|
||||
| 非显然约束 | 只解释"为什么必须这样"及违反后果,不解释"这行在干什么";load-bearing 的反直觉点必须写 | `# π̂ 必须 detach:否则学生通过抬高自身先验自我强化,训练塌缩` |
|
||||
| 非显然约束 | 只解释"为什么必须这样"及违反后果,不解释"这行在干什么";load-bearing 的反直觉点必须写 | `# π̂ 必须 detach:否则优化器会压低学生自身概率把乘子 π̂ 推向 0 以逃逸惩罚,teacher 否定的 chunk 最先塌缩` |
|
||||
| 差异标注 | 凡有意偏离论文或参考实现处,注明对方做法与我们的理由 | `# 参考实现(trainer:2205)对 chunk 内取 mean,论文式(8)为 sum,此处从论文` |
|
||||
|
||||
**类型与 shape**:
|
||||
|
||||
@@ -24,6 +24,19 @@ flowchart LR
|
||||
|
||||
$$\mathcal{L}_{\text{OmniOPD}}(\theta) = -\mathbb{E}_{\hat y\sim\pi_\theta}\Big[\sum_{c=1}^{M}\hat\pi^{(c)}_{\text{teacher}}\sum_{t\in c}\log\pi_\theta(y_t\mid x,y_{<t})\Big] + \beta\sum_{t\in\mathcal{U}} D_{KL}\big(\pi_{\text{ref}}\,\|\,\pi_\theta\big)$$
|
||||
|
||||
| 符号 | 含义 | 直观说法 |
|
||||
| --------------------------------- | ------------------------------------ | ------------------------------- |
|
||||
| $\pi_\theta$ | 学生模型($\theta$ 是它的参数,训练改的就是 $\theta$) | 正在被训练的 0.6B |
|
||||
| $\hat y \sim \pi_\theta$ | 轨迹是学生自己生成的 | “on-policy”三个字的全部含义 |
|
||||
| $\mathbb{E}[\cdot]$ | 期望 | 实践中 = 对 batch 里采样出的轨迹求平均,没有更多玄机 |
|
||||
| $c$,共 $M$ 个 | 被熵调度器选中的 chunk(各 $C=50$ 个 token) | 被“抽查”的 $M=10$ 段 |
|
||||
| $\hat\pi^{(c)}_{\text{teacher}}$ | 式(5)算出的贝叶斯估计,$\in [0,1]$ | 老师对这段的认可度打分 |
|
||||
| $\log \pi_\theta(y_t \mid \cdot)$ | 学生给自己当时生成的那个 token 的对数概率 | SFT 里最熟悉的那个量 |
|
||||
| $\mathcal{U}$ | 未被抽查的所有 token | 轨迹的绝大部分 |
|
||||
| $\pi_{\text{ref}}$ | 训练开始前学生的冻结副本 | “初始的自己” |
|
||||
| $\beta$ | 缰绳松紧 | 代码里的 `mc_kl_weight` |
|
||||
|
||||
|
||||
关键设计洞察(§4.1,Theorem 4.1):teacher 估计 π̂ 以**有界乘子** [0,1] 的身份乘在学生 score function 上,而不是像反向 KL 那样出现在分母/log 里——这从结构上消灭了标准 OPD 的梯度爆炸;而贝叶斯先验保证 π̂ ≥ α·π̄/(N+α) > 0,消灭了"teacher 全不匹配 ⇒ 梯度归零"的监督塌缩。
|
||||
|
||||
## 2. 参考实现的真实形态:一个 trainer,三代方法
|
||||
@@ -74,7 +87,7 @@ $$\bar\pi_\theta^{(c)} = \Big(\prod_{t\in c}\pi_\theta(y_t\mid\cdot)\Big)^{1/C}
|
||||
|
||||
代码在 `_compute_chunk_ebopd_loss` 内 L2194-2202:先验 `pi_bar = exp(mean(chunk_lps.detach()))`(几何均值,与式 4 严格一致),`pi_hat = (k + chunk_alpha·pi_bar)/(chunk_mc_samples + chunk_alpha)`,随后 clamp 到 [1e-8, 1] 并 detach。
|
||||
|
||||
> **detach 是命门**:先验和 π̂ 都必须切断梯度,否则学生会通过抬高自己的先验来自我强化(reward hacking 式塌缩)。配置里 `mc_nll_weight`(L246)的注释明确警告开启会塌缩,默认 0。这一点论文只隐含在"π̂ 是目标而非变量"里,代码把它变成了硬约束。
|
||||
> **detach 是命门**:先验和 π̂ 都必须切断梯度(L2196、L2201)。若留梯度通路,损失 π̂·|Σlog π_θ| 中 π̂ 也随 θ 可动,最速下降方向变成**压低**学生对自己 token 的概率、把乘子 π̂ 推向 0(p·ln(1/p)→0,指数快过对数)——在 teacher 全否定(k≈0)的 chunk 上损失可一路逃逸到 0,贝叶斯安全底 α·π̄ 被优化器亲手拆除,Theorem 4.1(a) 的梯度有界性也随之失效(多出的 ∇π̂ 项与惊讶度成正比)。相邻的另一个陷阱:`mc_nll_weight`(config L246-252)给非 MC 位置加自身 NLL 正则,帮助文本明确警告 "non-zero values cause self-reinforcement collapse"(无条件复读自己→熵塌缩),默认 0——两者是方向相反的两种自指失败。论文只隐含在"π̂ 是目标而非变量"里,代码把它变成了硬约束。
|
||||
|
||||
对应理论:Theorem 4.1(b) 下界 π̂ ≥ α·π̄/(N+α) > 0;4.1(c) 偏差-方差分解,α 是噪声-偏移旋钮;Theorem 4.2 证明 N=10 是方差收益的甜点。
|
||||
|
||||
@@ -146,3 +159,5 @@ KL 锚(`mc_kl_weight` = 论文的 β,config L284,**默认 0**)实现与
|
||||
4. 实现的 KL 锚与论文式(8)有哪三处差异?
|
||||
5. API teacher 路径为什么强制 char 级编辑距离?分叉点为什么要对齐词边界?
|
||||
6. `no_bayesian` 消融等价于论文里的哪个估计器?§4.1 预言它会怎么失败?
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user