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

14 KiB
Raw Blame History

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 上的文本级语义验证":

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_modeconfig L223)在 compute_losstrainer L2841-2847)处分派:

mode 内容 论文对应 学习价值
"standard" GKD 式白盒蒸馏:token 级 JSD/KLbeta 插值方向,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.py3212 行,全部算法)、distillation_config.py599 行配置)、distillation.py(179 行 CLI 入口,只覆盖 standard 模式的示例)。teacher 客户端在 ../generation/vllm_client.pyopenrouter_client.py)。

3. 逐概念映射(论文 → 代码)

3.1 SFT 基线(式 1

标准交叉熵。代码:_compute_sft_loss L2776。触发条件:无 teacher 且 lmbda=0L2855)。若配了 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_weightconfig 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. 强制追加终端 chunkLC 处(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-idn=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_dataloaderL818)→ _RepeatBatchDataLoader 把同一 collated batch 重复 gradient_accumulation_steps 次(免重复分词的性能 trick,L346)。
  2. _prepare_inputsL859)在窗口起点调 _fill_bufferL873):按 lmbda 掷硬币分 on/off-policy → on-policy 走 _generate_student_completionsL1070vLLM 或 HF generate)。
  3. training_stepL3063)→ compute_lossL2841)按 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. logL3133)聚合指标上报 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 预言它会怎么失败?