层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:
@@ -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.3:reduction 实义为 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 接线
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
Reference in New Issue
Block a user