层2/U1: 新增 DistillConfig(式(2) 白盒蒸馏参数)+ docs/03 表述打磨

configs.py(对应 docs/03 §5 U1):
- DistillConfig 自包含、不继承 SFTConfig;teacher_model 进 config、student 留脚本
- 三处刻意缺席: 无 teacher_completions_path/max_length/top_k(现场生成+全词表)
- 两温度分名: kl_temperature(散度 softmax)vs gen_temperature(on-policy 采样)
- 默认即式(2): beta=1 反向KL、温度1、纯采样 top_p=1;lr=1e-6(论文§5.1蒸馏)
- __post_init__ 8 分支构造即校验(gen_temperature>0 护 on-policy 语义)

docs/03:
- §2.3 β 三副面孔表: 记号统一 π、补 mode-seeking 对称、附全词表/稀疏双镜像说明
- §2.3 三条实现约定(温度/log域/batchmean)由一句话拆成可扫读表格
- §3 偏差清单上方补统领抉择原则

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-07-19 03:37:44 -04:00
parent db392d9c81
commit e42af5256f
2 changed files with 161 additions and 5 deletions
+16 -4
View File
@@ -40,13 +40,23 @@ on/off-policy 抽签**不在这里**——在 `_prepare_inputs`→`_fill_buffer`
### 2.3 generalized_jsd_lossDT:2408-2491)——β 的三副面孔
记号:$\pi_\theta$ = student$\pi_T$ = teacherKL 里"在前"的那个分布是被求期望的一方。
| β | 数学 | 语义 | 代码 |
|---|------|------|------|
| 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 |
| 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 |
细节:温度在 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,重构时按实义命名
> 行号说明:上表指向 **全词表** 路径(`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!(本章最大陷阱)
@@ -75,6 +85,8 @@ server 路径更粗:teacher 只回传 top-k logprobs 三张表(DT:2676-2679
## 3. 与式(2) 的偏差清单(默认配置下)
抉择原则(下表每一行都由它推出,而非逐条拍板):**式(2) 的算法本质忠于论文(on-policy + 反向 KL + 全词表分布对齐);参考实现的默认近似是为它的处境——API 传输 + 大 teacher——妥协出来的,换了我们的处境(本地同 tokenizer 的 4B teacher + 4×A800)就不继承;不改优化方向的表面差异,选对诊断/教学最有利的;一般性凡免费且未来有用则保留、凡昂贵且当前数据上空转则删除。** 一个反直觉推论:正因处境不同,我们回归论文本质反而比参考默认更贴式(2)(支持集那行)。逐行推导见下,"我们层 2"列即结论。
| 项 | 参考实现默认 | 严格式(2) | 我们层 2 |
|----|--------------|-----------|----------|
| 支持集 | top-1 稀疏 + 尾桶 | 全词表 | **全词表**`top_k=0` 等价;同 tokenizer 本地 teacher 使我们能比参考默认更贴论文) |