docs: 第三章——层 2 white-box OPD(式(2) 精读、standard 路径解剖、U1-U5 重构任务)
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,113 @@
|
|||||||
|
# 03 · 层 2:White-box OPD 基线(token 级反向 KL)
|
||||||
|
|
||||||
|
> 本章目标:吃透论文式(2) 的 on-policy 白盒蒸馏及其梯度爆炸脆弱性(§4.1),解剖参考实现 `distillation_mode="standard"` 路径,定出层 2 的重构任务。行号缩写:`DT:` = `references/ars-opd/trl/trl/experimental/distillation/distillation_trainer.py`,`TR:` = `references/ars-opd/train_distillation.py`,`CFG:` = `.../distillation/distillation_config.py`。
|
||||||
|
|
||||||
|
## 1. 论文侧:式(2) 精读
|
||||||
|
|
||||||
|
$$\mathcal{L}_{\text{OPD}} = \mathbb{E}_{y \sim \pi_\theta(\cdot|x)}\Big[\sum_t D_{KL}\big(\pi_\theta(\cdot|y_{<t},x)\,\|\,\pi_T(\cdot|y_{<t},x)\big)\Big]$$
|
||||||
|
|
||||||
|
三个成分逐个看:
|
||||||
|
|
||||||
|
| 成分 | 含义 | 为什么 |
|
||||||
|
|------|------|--------|
|
||||||
|
| $y \sim \pi_\theta$ | **on-policy**:轨迹由 student 自己采样 | 治 SFT 的曝光偏差——式(1) 只在 teacher 轨迹上监督,student 推理时一旦走出熟悉区域就没见过纠正信号;on-policy 让 teacher 在"student 实际会犯错的地方"给监督 |
|
||||||
|
| $D_{KL}(\pi_\theta\|\pi_T)$ | **反向 KL**(student 在前)| mode-seeking:student 容量小,与其平摊质量模仿 teacher 全分布(前向 KL 的 mode-covering),不如集中质量学好 teacher 的主模式 |
|
||||||
|
| 逐 token 求和 | token 级分布对齐 | 这就是"白盒":需要 teacher 每个位置的**完整 logits**,因此 teacher 必须本地可跑、且与 student **同 tokenizer**(词表逐位对应才能算 KL) |
|
||||||
|
|
||||||
|
**§4.1 梯度爆炸(本层的理论主课,OmniOPD 的出发点)**:反向 KL 对 student logit 的梯度含 $\log\frac{\pi_\theta(v)}{\pi_T(v)}$ 项。on-policy 下 student 会采到 teacher 认为极差的 token($\pi_T \to 0$),此时 log 比值 $\to \infty$,单个 token 的梯度可以炸掉整个 batch。第一章已给过一句话版本,本层用代码把它钉死:重构任务里有一个"梯度范数随 $\pi_T$ 衰减而暴涨"的单元测试(对应层 3 detach 测试的姊妹篇——一个证明旧方案为什么坏,一个证明新方案为什么稳)。
|
||||||
|
|
||||||
|
顺带记住对比锚点:层 5 的式(8) 用"有界乘子 π̂"替换这里的 log 比值,这正是两代方法的分水岭。
|
||||||
|
|
||||||
|
## 2. 参考实现解剖(standard 路径)
|
||||||
|
|
||||||
|
### 2.1 一个重要的事实先行
|
||||||
|
|
||||||
|
**训练脚本从不传本地 teacher**(TR:420-424 只传 model/args/dataset):仓库实际跑的 standard 蒸馏全走 **teacher server 路径**(环境变量 `TEACHER_URL` 触发,TR:329-330 → DT:2891-2920);**本地 teacher 路径(DT:2921-2944)只有直接构造 `DistillationTrainer(teacher_model=...)` 才会触发**。我们层 2 恰恰要走后者(4 卡 A800 放得下 4B teacher),所以两条路都要看懂,但以本地路为重构蓝本。
|
||||||
|
|
||||||
|
### 2.2 compute_loss 主流程(DT:2841-2946)
|
||||||
|
|
||||||
|
```
|
||||||
|
守卫: no_teacher and lmbda<1 → 显式报错 (DT:2869)
|
||||||
|
student 前向(带梯度) (DT:2883)
|
||||||
|
prompt_length 切齐 + [pl-1:-1] 移位 (DT:2887-2889) ← 与 SFT 完全同款,T4 已实现
|
||||||
|
teacher logits:
|
||||||
|
server 路 → top-k logprobs 传输 (DT:2898)
|
||||||
|
本地路 → teacher.eval() + no_grad 前向 (DT:2923, 2578-2594)
|
||||||
|
divergence: generalized_jsd_loss / 稀疏快路 (DT:2926-2942)
|
||||||
|
```
|
||||||
|
|
||||||
|
on/off-policy 抽签**不在这里**——在 `_prepare_inputs`→`_fill_buffer`(见 2.5)。
|
||||||
|
|
||||||
|
### 2.3 generalized_jsd_loss(DT:2408-2491)——β 的三副面孔
|
||||||
|
|
||||||
|
| β | 数学 | 语义 | 代码 |
|
||||||
|
|---|------|------|------|
|
||||||
|
| 0 | $KL(\pi_T\|\pi_\theta)$ | 前向,mode-covering | DT:150-151 |
|
||||||
|
| 1 | $KL(\pi_\theta\|\pi_T)$ | **反向,式(2) 用这个** | DT:152-153 |
|
||||||
|
| (0,1) | $\beta KL(p_T\|m)+(1{-}\beta)KL(p_\theta\|m)$,$m=(1{-}\beta)p_\theta{+}\beta p_T$ | JSD 插值 | DT:154-162 |
|
||||||
|
|
||||||
|
细节:温度在 softmax 前除进两侧 logits(DT:2439-2440);全程 log 域运算(logsumexp 混合、`clamp_min(tiny)` 防 log0,DT:143,156-159);reduction=`batchmean` 实为 **sum / 有效 token 数**(labels≠-100 先滤,DT:2386-2399)——名字叫 batchmean,实义是 per-token mean,重构时按实义命名。
|
||||||
|
|
||||||
|
### 2.4 默认配置不是全词表 KL!(本章最大陷阱)
|
||||||
|
|
||||||
|
默认 `loss_top_k=1, loss_add_tail=True`(CFG:353,368)→ 走 **top-1 稀疏快路**(DT:2507-2549):支持集 = {实际采样 token} ∪ {teacher top-1} + **尾桶**(第 K+1 个桶收纳截断外全部概率质量:$\log(1-\sum e^{\text{top}k})$,DT:108-118,防 top-1 时 loss 平凡为 0)。要严格对齐式(2) 的全词表反向 KL,必须 `loss_top_k=0`。top-k>1 时支持集依 β 取 teacher top-k / student top-k / 两者并集去重(DT:2449-2467)。
|
||||||
|
|
||||||
|
server 路径更粗:teacher 只回传 top-k logprobs 三张表(DT:2676-2679),β>0 时 CFG 强制 top-1(CFG:540-544)。这些截断全是**传输/显存工程妥协**(第一章 §3.7 讨论过),不是论文成分。
|
||||||
|
|
||||||
|
### 2.5 buffer 与 on-policy 生成(DT:846-932, 1071-1263)
|
||||||
|
|
||||||
|
- `_RepeatBatchDataLoader` 把同一 collated batch 重复 `gradient_accumulation_steps` 次(DT:346-371),`_fill_buffer` 按**切片级**伯努利抽签 on/off-policy(`random() <= lmbda`,DT:879,主进程抽签后广播)。
|
||||||
|
- **off-policy 切片用数据集自带 completion**(DT:893-894),不是 teacher 采样——SFT 数据混训,非蒸馏。
|
||||||
|
- on-policy 切片:vLLM colocate 生成(按 `vllm_sync_frequency` 同步权重,DT:1094-1106)或 `model.generate`(DT:1113-1168);生成结果重建 input_ids/labels 写回 buffer(DT:1170-1263,labels 只在 completion 段有效)。
|
||||||
|
- loss 前向是对已生成序列的 teacher-forcing(DT:2883)——"采样一次、前向算分布",GKD 标准做法。
|
||||||
|
|
||||||
|
### 2.6 本地 teacher 的基础设施(DT:504-586)
|
||||||
|
|
||||||
|
加载后 `accelerator.prepare_model(teacher, evaluation_mode=True)`(DT:584,随 DDP 每卡一份副本);同 tokenizer 校验比较 `get_vocab()`(DT:734-740),不匹配在 compute_loss 处显式报错(DT:2876);前向 `eval() + no_grad`(DT:2581-2586)。
|
||||||
|
|
||||||
|
### 2.7 顺带发现的坑(解剖副产物)
|
||||||
|
|
||||||
|
| 坑 | 位置 | 说明 |
|
||||||
|
|----|------|------|
|
||||||
|
| `lmbda=1 + no_teacher` 穿过守卫后在深处崩 | DT:2872 vs DT:2593 | 报错文案宣称合法,实际必崩——守卫条件写错 |
|
||||||
|
| off-policy ≠ teacher 蒸馏 | DT:893-894 | 语义上是"混 SFT",文档易误读 |
|
||||||
|
| `num_generations>1 且 lmbda<1` 会造重复样本 | CFG:593-596 | 官方注释自己承认 |
|
||||||
|
|
||||||
|
## 3. 与式(2) 的偏差清单(默认配置下)
|
||||||
|
|
||||||
|
| 项 | 参考实现默认 | 严格式(2) | 我们层 2 |
|
||||||
|
|----|--------------|-----------|----------|
|
||||||
|
| 支持集 | top-1 稀疏 + 尾桶 | 全词表 | **全词表**(`top_k=0` 等价;同 tokenizer 本地 teacher 使我们能比参考默认更贴论文) |
|
||||||
|
| on-policy 比例 | lmbda=1.0 ✓ | 纯 on-policy | lmbda=1(不实现混合抽签) |
|
||||||
|
| KL 方向 | beta=1.0 ✓ | 反向 | β 作为参数保留(前向/反向/JSD 同一公式,纯逻辑函数顺手覆盖,also 层 5 KL 锚要用前向) |
|
||||||
|
| 温度 | 1.0 ✓ | 1 | 1.0 |
|
||||||
|
| reduction | per-token mean | 论文 token 求和 | per-token mean(与 SFT loss 同尺度,才能对比曲线;差一个常数因子不改优化方向) |
|
||||||
|
|
||||||
|
## 4. 保留 / 替代 / 删除
|
||||||
|
|
||||||
|
| 决策 | 项目 |
|
||||||
|
|------|------|
|
||||||
|
| **保留** | "生成一次 + teacher-forcing 前向"结构;prompt 边界切齐/重掩码(直接复用 T4 的 `compute_prompt_length`);teacher `eval+no_grad`;同 tokenizer 校验(显式报错);β 语义与温度;per-token mean |
|
||||||
|
| **替代** | vLLM 学生生成 → `model.generate`(0.6B 生成不慢,省掉 colocate+权重同步整套复杂度;卡了吞吐再回来接 vLLM,触发条件记录于此);teacher server → 本地 Qwen3-4B HF 前向(全词表精确);buffer/RepeatBatchDataLoader/切片抽签 → 每个 batch 现场生成现场用(lmbda=1 下 buffer 是纯开销) |
|
||||||
|
| **删除** | top-k/尾桶/top-1 稀疏快路(全词表放得下就不近似;显存账见 §5)、`reverse_kl_top_1_mode`、teacher server 客户端、Liger、off-policy 混训、on/off-policy 指标群 |
|
||||||
|
|
||||||
|
## 5. 重构任务(Claude 写码、你精读提问)
|
||||||
|
|
||||||
|
| # | 任务 | 落点 | 备注 |
|
||||||
|
|---|------|------|------|
|
||||||
|
| U1 | `DistillConfig` dataclass(teacher 模型、β、温度、生成参数) | `ars_opd/configs.py` | 追加,不动 SFTConfig |
|
||||||
|
| U2 | 纯逻辑 divergence:全词表 masked token-KL/JSD(β 参数化)+ 单测 | `ars_opd/trainer.py`(纯张量函数,同 `sft_loss` 地位) | 单测含**梯度爆炸演示**:$\pi_T\to0$ 时梯度范数暴涨的断言(§4.1 的可执行版本) |
|
||||||
|
| U3 | collator 放开 prompt-only(返回 prompts/prompt_attention_mask 供生成) | `ars_opd/data.py` | 兑现 T3 预留的口子;SFT 路径行为不变(回归测试盯住) |
|
||||||
|
| U4 | `DistillTrainer`:生成 → teacher no_grad 前向 → divergence | `ars_opd/trainer.py` | on-policy 生成用 `model.generate`;tokenizer 一致性构造时校验 |
|
||||||
|
| U5 | 自包含脚本 | `scripts/train_whitebox.sh` + `.py` | teacher Qwen/Qwen3-4B;数据复用同一 DAPO 子集(prompt-only,无需 teacher 缓存) |
|
||||||
|
|
||||||
|
**显存账(A800-80G,验证 U 系列前算给自己看)**:全词表 logits (B,T,V) 是大头——B=4、T=2048、V≈151k 的 bf16 logits 单张 ≈2.5G,student+teacher 两份 + log_softmax 中间量 ≈15G 级,加 0.6B 训练全套 ~10G 与 4B teacher 推理副本 ~9G,B=4/T=2048 起步安全;生成长度先压 1024。参数进 `DistillConfig` 显式化。
|
||||||
|
|
||||||
|
**默认参数提案(可否决)**:β=1.0、temperature=1.0、lmbda 固定 1(不做参数)、`max_new_tokens=1024`、per_device batch 4、teacher `Qwen/Qwen3-4B`、数据同 seed 同 1k 子集(复用层 1 的抽取逻辑,不需要 teacher 解答缓存)。
|
||||||
|
|
||||||
|
## 6. 验证方式
|
||||||
|
|
||||||
|
1. **本地 CPU 单测**:divergence 手算对拍(V=5 玩具分布,β∈{0,1,0.5} 三点各一);β=1 与 β=0 的方向性断言(teacher 置信/弥散两种分布下 loss 排序);**梯度爆炸测试**:固定 student,teacher 对采样 token 的概率从 1e-1 衰减到 1e-6,断言梯度范数单调暴涨且超阈值——为层 5 的"有界乘子"对照埋桩。
|
||||||
|
2. **collator 回归**:放开 prompt-only 后,全部既有 SFT 测试必须原样通过。
|
||||||
|
3. **远程冒烟**:50 步,盯 KL loss 曲线与生成样本质量;预期能看到 loss 毛刺(梯度爆炸的实况)——这本身就是教学目标,截图留档给层 5 当对比。
|
||||||
|
4. 关账判据:白盒蒸馏 1k 子集跑通不 NaN(允许毛刺),生成质量肉眼不劣于 SFT 基线;接口回看完成。
|
||||||
Reference in New Issue
Block a user