层2/U2: 纯逻辑 token_divergence(式(2) 全词表 KL/JSD)+ 梯度爆炸演示单测

trainer.py(对应 docs/03 §5 U2,与 sft_loss 同为纯张量损失函数):
- token_divergence: 全词表精确 KL(β=1 反向=式(2) / β=0 前向 / (0,1) JSD)
- 只吃两组已对齐 logits + labels 掩码,不做移位(复用 T4 几何,留给 U4)
- 删参考实现 top-k/尾桶/nan_to_num(本地全词表恒有限);按 β 分支省一份 probs
- per-token mean 与 sft_loss 同尺度;全掩码/beta 越界显式报错

test_divergence.py(对应 docs/03 §6.1):
- 对拍 PyTorch torch.distributions.kl_divergence(独立 oracle,非同式自证)
- 方向性: 反向罚越界(mode-seeking) / 前向罚漏覆盖(mode-covering)
- §4.1 梯度爆炸: teacher 概率 1e-1→1e-6 时 student 梯度范数单调暴涨(层5对照桩)

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-07-19 04:00:25 -04:00
parent e42af5256f
commit de7f36828a
2 changed files with 211 additions and 0 deletions
+80
View File
@@ -91,6 +91,86 @@ def sft_loss(
return loss, num_valid
def token_divergence(
student_logits: torch.Tensor,
teacher_logits: torch.Tensor,
labels: torch.Tensor,
beta: float = 1.0,
temperature: float = 1.0,
) -> tuple[torch.Tensor, int]:
"""式(2)completion 位置上的 token 级(广义)KL 散度,全词表精确。
论文锚点:§3.1 式(2) L = E_{y~π_θ}[Σ_t KL(π_θ(·|y_<t,x) ‖ π_T(·|y_<t,x))]。
本函数只管"给定两组**已对齐**的 logits,算散度标量"——on-policy 采样(y~π_θ)
与移位对齐([pl-1:-1])由调用方(DistillTrainer, U4)负责,复用 sft_loss 同一套
切片几何。故这里不吃 input_ids:散度是分布对分布,不需要目标 token,labels
仅用于定位有效位置。
参数(B=batch,T=已移位对齐长度,V=词表大小):
- student_logits: (B, T, V)student 前向输出(带梯度)
- teacher_logits: (B, T, V)teacher no_grad 前向输出(无梯度)
- labels: (B, T)completion 位为 token id、其余为 IGNORE_INDEX;只做有效位掩码
- beta: KL 方向(docs/03 §2.3 三副面孔)。0=前向 KL(π_T‖π_θ)、1=反向 KL(π_θ‖π_T)=
式(2)、(0,1)=JSD 插值
- temperature: softmax 前除进两侧 logits 的温度(§2.3),软化/锐化分布
返回 (per-token mean 散度标量, 有效 token 数)。
差异标注:参考实现(DT:2408-2491)含 top-k 稀疏 + 尾桶快路,我们只保留全词表这
一条精确路径(docs/03 §4 删除清单:本地同 tokenizer teacher 放得下)。因全词表
log_softmax 对有限 logits 恒有限,也不需要参考在 -inf 支持集上的 nan_to_num 兜底。
"""
if not 0.0 <= beta <= 1.0:
raise ValueError(f"beta 必须在 [0,1]0=前向/1=反向/中间=JSD),收到 {beta}")
# 温度除进 logits、softmax 之前(§2.3):调分布形状,不是等比缩概率
student_logits = student_logits / temperature
teacher_logits = teacher_logits / temperature
# 全程 log 域(§2.3 数值稳定性)。log_probs 两侧都要;probs 按分支只算用得上的
# 那一份——(B,T,V) 在真实规模下每份 ~2.5G(§5 显存账),默认 β=1 热路径只需
# student_probs,不materialize teacher_probs
student_log_probs = F.log_softmax(student_logits, dim=-1) # (B, T, V)
teacher_log_probs = F.log_softmax(teacher_logits, dim=-1) # (B, T, V)
if beta == 1.0:
# 反向 KL(π_θ‖π_T) = Σ_v π_θ (log π_θ log π_T) —— 式(2)
# 非显然约束(§4.1 梯度爆炸源):对 student logit 的梯度含 π_θ·log(π_θ/π_T),
# on-policy 采到 teacher 眼中烂 token(π_T→0)时 log 比值→∞,单 token 梯度可
# 炸掉整个 batch。这正是层 2 要亲眼观察、层 5 用有界乘子 π̂ 替换的病灶
per_token = (
student_log_probs.exp() * (student_log_probs - teacher_log_probs)
).sum(-1) # (B, T, V) -> (B, T)
elif beta == 0.0:
# 前向 KL(π_T‖π_θ) = Σ_v π_T (log π_T log π_θ)
per_token = (
teacher_log_probs.exp() * (teacher_log_probs - student_log_probs)
).sum(-1)
else:
# JSD 插值:m = (1−β)π_θ + β π_T;β·KL(π_T‖m) + (1−β)·KL(π_θ‖m)
student_probs = student_log_probs.exp()
teacher_probs = teacher_log_probs.exp()
mixture = (1.0 - beta) * student_probs + beta * teacher_probs
# clamp_min(tiny) 防 log0(§2.3):混合概率理论上恒正,此处是浮点下溢兜底
log_mixture = mixture.clamp_min(torch.finfo(mixture.dtype).tiny).log()
kl_teacher = (teacher_probs * (teacher_log_probs - log_mixture)).sum(-1)
kl_student = (student_probs * (student_log_probs - log_mixture)).sum(-1)
per_token = beta * kl_teacher + (1.0 - beta) * kl_student
# 掩码 + per-token mean(docs/03 §2.3reduction 实义为 sum/有效token数,
# 与层 1 sft_loss 同尺度,两条 loss 曲线才可比)
mask = labels != IGNORE_INDEX # (B, T)
num_valid = int(mask.sum().item())
if num_valid == 0:
# 与 sft_loss 同纪律:走到这里全被掩码只可能是数据/生成坏了,显式报错不静默
raise ValueError(
"本 batch 没有任何有效 completion token(全被 -100 掩码)。"
"检查 on-policy 生成是否产出了空 completion。"
)
loss = per_token[mask].sum() / num_valid
return loss, num_valid
# ---------------------------------------------------------------------------
# Trainer 接线
# ---------------------------------------------------------------------------