docs: 第一章——论文精读与参考实现全景映射(含实现与论文差异清单)

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
2026-07-17 11:29:20 -04:00
parent 6dc13314d8
commit 8647c5a89d
2 changed files with 149 additions and 1 deletions
+148
View File
@@ -0,0 +1,148 @@
# 01 · 论文精读与参考实现全景映射
> 本章目标:读完后你应该能 (a) 复述 OmniOPD 的完整方法链条;(b) 对论文任一公式说出它在参考实现的哪个文件哪一行;(c) 说出实现与论文的全部差异。
> 参考实现路径统称 `distillation/` = `references/ars-opd/trl/trl/experimental/distillation/`。行号基于当前快照。
## 1. 方法总览
OmniOPD 要解决标准 OPD 的两个耦合缺陷:① 需要 teacher 的 token 级 logits(商业模型不给);② token 级 logit 匹配信号本身脆弱(依赖师生 plausible-token 集合的狭窄交集,tokenizer 不匹配、重复循环会污染梯度)。解法是把监督信号从"每个 token 的 logit 分布"换成"少数关键 chunk 上的文本级语义验证"
```mermaid
flowchart LR
A[学生在线生成轨迹 y] --> B[逐 token 熵 H_t<br/>式6]
B --> C[选 M 个高熵分叉点<br/>各取 C-token chunk 式7]
C --> D[对每个 chunkteacher 按前缀<br/>采 N 条 rollout]
D --> E[语义相似度 φ 求和<br/>k_sem 式3]
E --> F[Dirichlet 平滑得 π̂<br/>式4-5]
F --> G[chunk 损失:-π̂·Σlog π_θ<br/>式8 左项]
A --> H[未审计 token<br/>对冻结初始模型的 KL 锚 式8 右项]
G --> I[总损失]
H --> I
```
总目标(式 8):
$$\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)$$
关键设计洞察(§4.1Theorem 4.1):teacher 估计 π̂ 以**有界乘子** [0,1] 的身份乘在学生 score function 上,而不是像反向 KL 那样出现在分母/log 里——这从结构上消灭了标准 OPD 的梯度爆炸;而贝叶斯先验保证 π̂ ≥ α·π̄/(N+α) > 0,消灭了"teacher 全不匹配 ⇒ 梯度归零"的监督塌缩。
## 2. 参考实现的真实形态:一个 trainer,三代方法
`DistillationTrainer` 实际叠合了方法演化的三个阶段,由 `distillation_mode`config L223)在 `compute_loss`trainer L2841-2847)处分派:
| mode | 内容 | 论文对应 | 学习价值 |
|------|------|----------|----------|
| `"standard"` | GKD 式白盒蒸馏:token 级 JSD/KL`beta` 插值方向,`lmbda` 控制 on/off-policy 比例,`loss_top_k` 截断词表 | §3.1 式(2),被超越的 baseline | 层 2 要重建的对象 |
| `"ebopd"` | **token 级** MC:在高熵 token 处向 teacher 采 1-token rollout 数匹配次数 | §3.2.1 开头被否掉的 naive 方案(查询量 O(T)、精确 token 匹配稀疏) | 方法演化的"中间化石",论文没写但代码留着 |
| `"chunk_ebopd"` | **chunk 级** MC + 贝叶斯 + 熵调度 + KL 锚 | §3.2 完整 OmniOPD | 层 3-5 要重建的对象 |
> 命名注记:代码内部叫 EB-OPDentropy-based),论文成稿改名 OmniOPD;类元数据(L390)还引着 Agarwal 2024 的老 OPD 论文,属于历史遗留。
三个文件分工:`distillation_trainer.py`3212 行,全部算法)、`distillation_config.py`599 行配置)、`distillation.py`(179 行 CLI 入口,只覆盖 standard 模式的示例)。teacher 客户端在 `../generation/``vllm_client.py``openrouter_client.py`)。
## 3. 逐概念映射(论文 → 代码)
### 3.1 SFT 基线(式 1
标准交叉熵。代码:`_compute_sft_loss` L2776。触发条件:无 teacher 且 `lmbda=0`L2855)。若配了 teacher server,则由 `_generate_teacher_completions`(L934)先生成离线数据——这里有 OpenRouter 的 JSONL 落盘缓存(L945-1068,防重复付费)。
### 3.2 白盒 OPD 基线(式 2
$$\mathcal{L}_{\text{OPD}} = \mathbb{E}\Big[\sum_t D_{KL}\big(\pi_\theta(\cdot\mid x,\hat y_{<t})\,\|\,\pi^*_{\text{teacher}}(\cdot\mid x,\hat y_{<t})\big)\Big]$$
| 组件 | 代码 | 说明 |
|------|------|------|
| 通用 JSD 损失 | `generalized_jsd_loss` L2408、`_jsd_divergence` L121 | `beta=0` 前向 KL、`beta=1` 反向 KL、中间值 JSD |
| top-k 支持截断 | `loss_top_k` config L352、tail bucket `_add_tail_bucket` L108 | 只在 K 个 token 上算散度 + 一个"尾桶"吸收剩余概率质量——**论文没有**,是服务器传输 logprobs 的工程妥协 |
| teacher logits 三来源 | 本地 `_get_teacher_logits` L2578server `_get_teacher_token_logprobs_from_server` L2596 | server 只能拿 top-k logprobs,故必须截断 |
| on/off-policy 混合 | `lmbda` 掷硬币 L878-882 | `lmbda=1` 纯 on-policy |
### 3.3 语义相似度 φ(式 3)
$$k_{\text{sem}}^{(c)} = \sum_{i=1}^{N}\phi(y_c,\,y^{(i)}_{\text{teacher}})$$
| 实现 | 代码 | 备注 |
|------|------|------|
| ROUGE-1 F1 | `_compute_rouge1` L1669 | unigram 集合交并 → 2PR/(P+R+1e-8),白空格分词 |
| 编辑距离相似度 | `_compute_edit_similarity` L1681 | 双行 Levenshtein DP → 1 dist/max(m,n) |
| 选择开关 | `chunk_similarity` config L298,默认 rouge1 | 本地 vLLM 路径:edit 用 token-id 序列、rouge 用解码文本(L1840-1852 |
| **API 路径强制 char 级 edit** | L2049,理由在 L1897-1901 | 因为 `text.split()` 对中文失效——**论文没提的 CJK 处理** |
### 3.4 贝叶斯平滑(式 4-5)——方法的心脏
$$\bar\pi_\theta^{(c)} = \Big(\prod_{t\in c}\pi_\theta(y_t\mid\cdot)\Big)^{1/C} \qquad \hat\pi^{(c)}_{\text{teacher}} = \frac{k^{(c)}_{\text{sem}} + \alpha\,\bar\pi^{(c)}_\theta}{N+\alpha}$$
代码在 `_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。这一点论文只隐含在"π̂ 是目标而非变量"里,代码把它变成了硬约束。
对应理论:Theorem 4.1(b) 下界 π̂ ≥ α·π̄/(N+α) > 0;4.1(c) 偏差-方差分解,α 是噪声-偏移旋钮;Theorem 4.2 证明 N=10 是方差收益的甜点。
### 3.5 Peak-entropy 调度器(式 6-7
论文:取轨迹中熵最高的 M 个 token 作锚点(式 7 的 top-M),各扩成 C-token chunk,重叠则合并/重采。代码 `_get_chunk_fork_positions` L1700-1762 的实际算法更具体:
1. 无效位、padding 位熵置 −1(L1716);保留尾部 `[L2C, L]` 不参与选择(L1722)。
2. 贪心循环(L1747-1756):取 `argmax`,然后把 `±chunk_min_distance`(默认 50 = C)窗口内熵置 −1——用**间距抑制**代替论文的"重叠合并",天然保证 chunk 互不重叠。
3. **强制追加终端 chunk**`LC` 处(L1758-1760)——轨迹结尾(最终答案)永远被审计。论文正文没有此规则,但消融开关 `no_terminal_chunk` 暴露了它。
token 级 ebopd 模式用的是另一套:分位数阈值 `entropy_percentile``_get_high_entropy_mask` L1293,跨进程全局 quantile)——可对比学习两种选择策略。
### 3.6 Teacher MC rollout(式 3 的采样端)
| 路径 | 代码 | 要点 |
|------|------|------|
| 本地 vLLM teacher | `_mc_sample_teacher_chunks` L1764 | 前缀 = prompt + 完成到分叉点的 token-id`n=N, max_tokens=C` 一次采齐 |
| API teacherOpenRouter/OpenAI 兼容) | `_mc_sample_teacher_chunks_api` L1874 | 文本前缀;`max_tokens=C+16`;两种续写方式:assistant-prefill 或 prompt-continuationL2004-2022),按模型白名单 `_PREFILL_SUPPORTED_MODELS` L1659 分派 |
| 分叉点词边界对齐 | `_snap_fork_to_word_boundary` L1905-1926 | API 路径把分叉点前挪 ≤5 token 到 ASCII 词边界,CJK 视为天然边界 |
| 空回复补偿 | L2052-2057 | 429 限流导致的空 rollout 剔除后按均值重放大 k_score,保持目标尺度 |
| 限流与重试 | `openrouter_client.py`:并发信号量 4,指数退避 4 次 | chunk MC **不落盘缓存**(依赖服务商前缀缓存),只有 SFT 生成路径有 JSONL 缓存 |
### 3.7 总损失与 trust-region 锚(式 8
chunk 项:`chunk_loss = -pi_hat * chunk_lps.mean()`L2205),全 chunk 取 `mean` 聚合(L2209)。注意实现对 chunk 内 token 用的是 **mean** 而论文式 8 是 **sum**——差一个 1/C 常数,被学习率吸收。
KL 锚(`mc_kl_weight` = 论文的 β,config L284,**默认 0**)实现与论文有三处实质差异,重构时必须决策:
| # | 论文 | 实现 | 位置 |
|---|------|------|------|
| a | KL 只施加于未审计集合 U | 施加于**全部** completion token(含已审计 chunk | L2213-2306 |
| b | D_KL(π_ref ∥ π_θ) 全词表 | server 路径:前向 KL 但只在 ref 的 top-10 logprobs 支持上 | L2249-2251 |
| c | 同一公式 | 本地 ref 回退路径:Schulman k3 估计器(reverse 方向),CPU↔GPU 搬运 + 分段前向省显存 | L2271-2306 |
### 3.8 论文没写、实现里有的稳定器与消融
| 机制 | 代码 | 作用 |
|------|------|------|
| 长度权重 `sigmoid(\|a_len\|+2)` | L2325-2342,冻结的跨 batch 长度锚点 | 降权长度异常的轨迹,防长度投机 |
| 熵地板 hinge | L2308-2322`_entropy_anchor` 一次性锚定 | 防熵塌缩(人为缩短轨迹逃避审计) |
| 消融 `no_terminal_chunk` | L1720-1731, L1759 | 关终端 chunk |
| 消融 `no_bayesian` | L2188-2193 + 强制均匀选 chunk | 直接用 k/N 当权重(即 §4.1 会塌缩的频率派估计) |
| 消融 `uniform_chunks` | L1734-1744 | 随机选 chunk 但保留贝叶斯 π̂ |
## 4. 一次训练步的完整流程
1. `get_train_dataloader`L818)→ `_RepeatBatchDataLoader` 把同一 collated batch 重复 `gradient_accumulation_steps` 次(免重复分词的性能 trick,L346)。
2. `_prepare_inputs`L859)在窗口起点调 `_fill_buffer`L873):按 `lmbda` 掷硬币分 on/off-policy → on-policy 走 `_generate_student_completions`L1070vLLM 或 HF generate)。
3. `training_step`L3063)→ `compute_loss`L2841)按 `distillation_mode` 分派。
4. chunk_ebopd 链:学生前向(L2124)→ logits 切片 `[:, prompt_length-1:-1]`L2137,注意 off-by-one 对齐)→ 熵(L2146)→ 选 chunkL2147)→ MC 采样(L2152-2164)→ 贝叶斯 chunk 损失(L2175-2211)→ ref-KLL2216)→ 熵地板(L2311)→ 长度权重(L2325)。
5. `log`L3133)聚合指标上报 W&B。
## 5. 重构启示(映射到我们的模块)
| 我们的模块 | 吸收参考实现的 | 舍弃/待议 |
|------------|----------------|-----------|
| `similarity.py` | rouge1 / edit 两个纯函数;CJK→char 级的教训 | — |
| `estimator.py` | 式4-5 + clamp + **detach 硬约束**;空样本重放大 | token 级 ebopd 路径不重建(只读懂) |
| `chunking.py` | 贪心 argmax + 间距抑制 + 终端 chunk;三个消融开关 | percentile 全局分位数(那是 ebopd 的) |
| `teacher.py` | OpenAI 兼容客户端 + 并发/退避 + prefill/continuation 双路 + 词边界对齐 | OpenRouter 专属逻辑;`hash()` 缓存键(不稳定,须换 sha256);chunk MC 也应加落盘缓存(原版没有) |
| `trainer.py` | 式8 编排 + 长度权重 + 熵地板 | Liger/ZeRO-3、`_RepeatBatchDataLoader` 视规模再说;KL 锚按论文改回"仅未审计 token"还是保持全 token?(层 5 实验决定) |
## 6. 自查问题(学完本章应能回答)
1. 为什么式(8)把 π̂ 放在学生 log-likelihood 外面做乘子,而不是像式(2)放进 KL?这如何同时解决梯度爆炸与监督塌缩?
2. α→0 和 α→∞ 各退化成什么?为什么 N=10 就够(Theorem 4.2 的论证链)?
3. π̂ 若不 detach 会发生什么?哪条配置项的注释印证了这一点?
4. 实现的 KL 锚与论文式(8)有哪三处差异?
5. API teacher 路径为什么强制 char 级编辑距离?分叉点为什么要对齐词边界?
6. `no_bayesian` 消融等价于论文里的哪个估计器?§4.1 预言它会怎么失败?