Files
ars-opd-rebuild/docs/03-whitebox-opd.md
T

114 lines
11 KiB
Markdown
Raw 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.
# 03 · 层 2White-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-seekingstudent 容量小,与其平摊质量模仿 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_lossDT: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 前除进两侧 logitsDT:2439-2440);全程 log 域运算(logsumexp 混合、`clamp_min(tiny)` 防 log0DT: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-1CFG: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 的轨迹上做 KL 蒸馏**DT:893-894 原样保留数据轨迹;损失仍是 KL,不是交叉熵)——GKD 的 λ 插值本义:λ=0 离线蒸馏、λ=1 纯 on-policy。两个易误解处:轨迹不是 teacher 现场采样的;DAPO prompt-only 下这些切片 labels 全 -100KL 被掩码归零 = **静默空转的算力浪费**(唯一例外:`lmbda=0`+server 触发 teacher 生成,DT:901-906,即层 1 SFT 的数据来源)。
- on-policy 切片:vLLM colocate 生成(按 `vllm_sync_frequency` 同步权重,DT:1094-1106)或 `model.generate`DT:1113-1168);生成结果重建 input_ids/labels 写回 bufferDT:1170-1263labels 只在 completion 段有效)。
- loss 前向是对已生成序列的 teacher-forcingDT: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 | 轨迹来自数据集固有 completion(损失仍是 KL);prompt-only 数据下整个切片被掩码归零,静默空转 |
| `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` dataclassteacher 模型、β、温度、生成参数) | `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.5Gstudent+teacher 两份 + log_softmax 中间量 ≈15G 级,加 0.6B 训练全套 ~10G 与 4B teacher 推理副本 ~9GB=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 排序);**梯度爆炸测试**:固定 studentteacher 对采样 token 的概率从 1e-1 衰减到 1e-6,断言梯度范数单调暴涨且超阈值——为层 5 的"有界乘子"对照埋桩。
2. **collator 回归**:放开 prompt-only 后,全部既有 SFT 测试必须原样通过。
3. **远程冒烟**:50 步,盯 KL loss 曲线与生成样本质量;预期能看到 loss 毛刺(梯度爆炸的实况)——这本身就是教学目标,截图留档给层 5 当对比。
4. 关账判据:白盒蒸馏 1k 子集跑通不 NaN(允许毛刺),生成质量肉眼不劣于 SFT 基线;接口回看完成。