Files
iomgaa fac8e0dcc5 docs: 层 2 关账——§6.1 远程实证勘误"预期见毛刺" + roadmap 存档点 + 接口回看
- docs/03 §6.1:两次远程跑(sanity/noclip)均平稳,勘误当初"预期毛刺"的错误
  预测;三条实证结论(爆炸真机制但本区间高度阻尼、grad_norm 是裁剪前值、层 5
  真正杀手锏是 logit-free)
- roadmap 存档点:层 2  关账,判据全过、关键发现、接口回看(全判"深",
  三条不返工小注);章节索引 03 标已关账;下一步层 3

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-19 09:30:45 -04:00

141 lines
15 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.
# 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)——β 的三副面孔
记号:$\pi_\theta$ = student$\pi_T$ = teacherKL 里"在前"的那个分布是被求期望的一方。
| β | 数学 | 语义 | 代码 |
|---|------|------|------|
| 0 | $KL(\pi_T\|\pi_\theta)$ | 前向,**mode-covering**teacher 在前,student 被迫摊平质量去覆盖 teacher 的全分布 | DT:150-151 |
| 1 | $KL(\pi_\theta\|\pi_T)$ | 反向,**mode-seeking****式(2) 用这个**student 在前,集中质量学 teacher 主模式 | DT:152-153 |
| (0,1) | $\beta KL(\pi_T\|m)+(1{-}\beta)KL(\pi_\theta\|m)$$m=(1{-}\beta)\pi_\theta{+}\beta\pi_T$ | JSD 插值(β 同时是混合权重与两项权重) | DT:154-162 |
> 行号说明:上表指向 **全词表** 路径(`F.kl_div`,我们 U2 走这条)。参考实现**默认**走 top-1 稀疏(§2.4),对应 DT:133-148 的 masked 镜像分支——同样三支 β、同样数学,只是在截断支持集上手算而非调 `F.kl_div`。
另有三条与 β 语义无关、但读代码时容易卡住的实现约定。它们各自独立,只是恰好都在这个函数里;前两条是正确性/稳定性刚需(我们保留),第三条是历史包袱(我们纠名):
| 实现约定 | 位置 | 是什么 / 为什么 load-bearing | 我们 U2 |
|----------|------|------------------------------|---------|
| 温度除进 logits | DT:2439-2440 | `logits / τ` 必须在 softmax **之前**做——这是在调分布形状(升温 τ>1 放大尾部的"暗知识"排序),不是等比缩概率。`softmax(z/τ) ≠ softmax(z)/τ`,位置错了就不再是合法分布 | 保留(默认 τ=1,此步为恒等) |
| 全程 log 域运算 | DT:143,156-159 | 15 万词表下单个概率小到 1e-8,而 KL 全是乘除,直接算会下溢成 0 → NaN。对策:全程存 log-prob(乘变加、除变减)。两个衍生 trick:混合分布 $m$ 的**加法**在 log 域要用 `logsumexp`(log 里的加法天然是乘法);masked 位概率为 0,取 log 前先 `clamp_min(tiny)``log(0)=-inf` | 保留(仅全词表这一路径) |
| `batchmean` 名不副实 | DT:2386-2399 | 先滤 `labels≠-100` 留下 completion 位,再 `jsd.sum() / 有效 token 数`——实义是 **per-token mean**,不是 PyTorch `batchmean` 那个"÷ 序列条数"。这样量纲与 SFT 的 per-token 交叉熵一致,两条 loss 曲线才可比 | 纠名为 `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) 的算法本质忠于论文(on-policy + 反向 KL + 全词表分布对齐);参考实现的默认近似是为它的处境——API 传输 + 大 teacher——妥协出来的,换了我们的处境(本地同 tokenizer 的 4B teacher + 4×A800)就不继承;不改优化方向的表面差异,选对诊断/教学最有利的;一般性凡免费且未来有用则保留、凡昂贵且当前数据上空转则删除。** 一个反直觉推论:正因处境不同,我们回归论文本质反而比参考默认更贴式(2)(支持集那行)。逐行推导见下,"我们层 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. **远程冒烟**:盯 KL loss 曲线、`grad_norm``distill/num_gen_tokens_per_step`
4. 关账判据:白盒蒸馏跑通不 NaN,生成不塌空,接口回看完成。
### 6.1 远程实证(2026-07-19,两次跑,勘误当初的"预期见毛刺")
**当初预测错了**docs 原写"预期能看到 loss 毛刺(梯度爆炸实况)"。实跑**没有毛刺**,两次都平稳。诚实记录 + 解释:
| 跑 | 配置 | loss | grad_norm |
|----|------|------|-----------|
| sanity | 裁到 1.0、lr 1e-6、50 步 | 平滑 0.35→0.22 | 14.42 → ~2(单调降) |
| noclip | ≈关裁剪、lr 5e-6、15 步 | 平滑 0.35→0.21 | 14.42 → ~2(无尖峰,且更快收敛) |
三条实证结论:
- **爆炸是真机制、但本区间高度阻尼**。§4.1 在单测里坐实(单个 π_T→1e-6 的 token 梯度暴涨),但真实训练里:① **同门 teacher**Qwen3 0.6B↔4B)使 student 很少采到 teacher 真恨的 token;② **per-token mean 把每步 ~4000 token 的梯度尖峰摊平**(单测看单 token 机制,真实看上千 token 平均后果)。故 batch 级 grad_norm 峰值只 ~14(比健康 ~2 高 5-7 倍,但远非几十上百),且只降不升。
- **关键坑:`grad_norm` 日志是裁剪前值**。sanity 的 14→2 那串本身就是爆炸证据,只是 HF 默认 `max_grad_norm=1.0` 把**步长**裁掉了、loss 才平滑——这个静默稳定器现已提进 `DistillConfig`(见其注释)。noclip 关掉它,grad_norm 曲线几乎不变(第 1 步两跑完全相同=14.42,验证确定性),但大步长反而**加速收敛**、仍不炸。
- **重构层 5 动机的认知**:白盒的"同 tokenizer"硬约束把你锁在相对温和的区间(换跨家族 teacher 会先被词表校验拦下),所以层 5 有界乘子 π̂ 的真正杀手锏不在"防这个温和爆炸",而在 **logit-free**(teacher 只给文本、拿不到 logits,白盒根本跑不了)。爆炸的干净见证留在 U2 单测。