Files
ars-opd-rebuild/docs/01-paper-code-map.md

164 lines
14 KiB
Markdown
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# 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)$$
| 符号 | 含义 | 直观说法 |
| --------------------------------- | ------------------------------------ | ------------------------------- |
| $\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.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 是命门**:先验和 π̂ 都必须切断梯度(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 是方差收益的甜点。
### 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 预言它会怎么失败?