From 3c1530a1dbf4d22ec5457e7238b59654f09d8808 Mon Sep 17 00:00:00 2001 From: iomgaa Date: Sat, 18 Jul 2026 08:34:29 -0400 Subject: [PATCH] =?UTF-8?q?docs:=20=E7=AC=AC=E4=B8=89=E7=AB=A0=E2=80=94?= =?UTF-8?q?=E2=80=94=E5=B1=82=202=20white-box=20OPD=EF=BC=88=E5=BC=8F(2)?= =?UTF-8?q?=20=E7=B2=BE=E8=AF=BB=E3=80=81standard=20=E8=B7=AF=E5=BE=84?= =?UTF-8?q?=E8=A7=A3=E5=89=96=E3=80=81U1-U5=20=E9=87=8D=E6=9E=84=E4=BB=BB?= =?UTF-8?q?=E5=8A=A1=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-Authored-By: Claude Fable 5 --- docs/03-whitebox-opd.md | 113 ++++++++++++++++++++++++++++++++++++++++ 1 file changed, 113 insertions(+) create mode 100644 docs/03-whitebox-opd.md diff --git a/docs/03-whitebox-opd.md b/docs/03-whitebox-opd.md new file mode 100644 index 0000000..a54a971 --- /dev/null +++ b/docs/03-whitebox-opd.md @@ -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_{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 基线;接口回看完成。