层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
|
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 接线
|
# Trainer 接线
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|||||||
@@ -0,0 +1,131 @@
|
|||||||
|
"""层 2 / U2:token 级散度单测(docs/03 §2、§4.1、§6.1)。
|
||||||
|
|
||||||
|
只测纯张量函数 token_divergence;DistillTrainer 类由远程冒烟验证。
|
||||||
|
三块:
|
||||||
|
1. 对拍 PyTorch 自带 KL(独立 oracle,非同式自证);
|
||||||
|
2. 反向/前向 KL 的方向性(mode-seeking vs mode-covering);
|
||||||
|
3. §4.1 梯度爆炸演示——teacher 概率趋 0 时 student 梯度暴涨(层 5 有界乘子的对照桩)。
|
||||||
|
|
||||||
|
构造技巧:softmax(log p) = p(p 已归一),故用 `probs.log()` 当 logits 即可精确
|
||||||
|
控制两侧分布,让手算/对拍成为可能。
|
||||||
|
"""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import torch
|
||||||
|
from torch.distributions import Categorical, kl_divergence
|
||||||
|
|
||||||
|
from ars_opd.data import IGNORE_INDEX
|
||||||
|
from ars_opd.trainer import token_divergence
|
||||||
|
|
||||||
|
V = 5 # 玩具词表
|
||||||
|
|
||||||
|
|
||||||
|
def logits_of(probs: list[float]) -> torch.Tensor:
|
||||||
|
"""概率向量 -> (1, 1, V) logits,使 log_softmax 后精确还原该分布。"""
|
||||||
|
return torch.tensor(probs).log().reshape(1, 1, V)
|
||||||
|
|
||||||
|
|
||||||
|
ONE_VALID = torch.zeros(1, 1, dtype=torch.long) # 单个有效 token(id 0 ≠ -100)
|
||||||
|
|
||||||
|
|
||||||
|
# ---- 1. 对拍 PyTorch KL ----
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("beta", [0.0, 1.0, 0.5])
|
||||||
|
def test_散度对拍pytorch_kl(beta):
|
||||||
|
p_s = [0.10, 0.20, 0.30, 0.25, 0.15]
|
||||||
|
p_t = [0.05, 0.05, 0.40, 0.40, 0.10]
|
||||||
|
loss, num = token_divergence(logits_of(p_s), logits_of(p_t), ONE_VALID, beta=beta)
|
||||||
|
assert num == 1
|
||||||
|
|
||||||
|
cs, ct = Categorical(torch.tensor(p_s)), Categorical(torch.tensor(p_t))
|
||||||
|
if beta == 1.0: # 反向 KL(π_θ‖π_T)
|
||||||
|
oracle = kl_divergence(cs, ct)
|
||||||
|
elif beta == 0.0: # 前向 KL(π_T‖π_θ)
|
||||||
|
oracle = kl_divergence(ct, cs)
|
||||||
|
else: # JSD:对混合分布的两支 KL 加权
|
||||||
|
m = Categorical((1 - beta) * torch.tensor(p_s) + beta * torch.tensor(p_t))
|
||||||
|
oracle = beta * kl_divergence(ct, m) + (1 - beta) * kl_divergence(cs, m)
|
||||||
|
assert torch.allclose(loss, oracle, atol=1e-6)
|
||||||
|
|
||||||
|
|
||||||
|
def test_同分布散度为零():
|
||||||
|
p = [0.1, 0.2, 0.3, 0.25, 0.15]
|
||||||
|
for beta in (0.0, 1.0, 0.5):
|
||||||
|
loss, _ = token_divergence(logits_of(p), logits_of(p), ONE_VALID, beta=beta)
|
||||||
|
assert torch.allclose(loss, torch.zeros(()), atol=1e-6)
|
||||||
|
|
||||||
|
|
||||||
|
# ---- 2. 方向性:反向罚"越界",前向罚"漏覆盖" ----
|
||||||
|
|
||||||
|
|
||||||
|
def test_kl方向性():
|
||||||
|
peaked = [0.90, 0.025, 0.025, 0.025, 0.025]
|
||||||
|
diffuse = [1 / V] * V
|
||||||
|
|
||||||
|
# 情形 A:teacher 尖、student 弥散——student 把质量放到 teacher≈0 处。
|
||||||
|
# 反向 KL(π_θ‖π_T) 因 log(π_θ/π_T) 在越界 token 上爆大而重罚;前向相对轻。
|
||||||
|
rev_A = token_divergence(
|
||||||
|
logits_of(diffuse), logits_of(peaked), ONE_VALID, beta=1.0
|
||||||
|
)[0]
|
||||||
|
fwd_A = token_divergence(
|
||||||
|
logits_of(diffuse), logits_of(peaked), ONE_VALID, beta=0.0
|
||||||
|
)[0]
|
||||||
|
assert rev_A > fwd_A # 反向惩罚 student 越出 teacher 支持集(mode-seeking)
|
||||||
|
|
||||||
|
# 情形 B:student 尖、teacher 弥散——teacher 的质量落在 student≈0 处。
|
||||||
|
# 前向 KL(π_T‖π_θ) 重罚"漏覆盖";反向相对轻。
|
||||||
|
rev_B = token_divergence(
|
||||||
|
logits_of(peaked), logits_of(diffuse), ONE_VALID, beta=1.0
|
||||||
|
)[0]
|
||||||
|
fwd_B = token_divergence(
|
||||||
|
logits_of(peaked), logits_of(diffuse), ONE_VALID, beta=0.0
|
||||||
|
)[0]
|
||||||
|
assert fwd_B > rev_B # 前向惩罚 student 没覆盖 teacher 的质量(mode-covering)
|
||||||
|
|
||||||
|
|
||||||
|
def test_温度升高软化分布降低反向kl():
|
||||||
|
# student 与 teacher 都尖但尖在不同 token;升温软化两侧 → 反向 KL 下降
|
||||||
|
s, t = [0.90, 0.025, 0.025, 0.025, 0.025], [0.025, 0.90, 0.025, 0.025, 0.025]
|
||||||
|
cold = token_divergence(logits_of(s), logits_of(t), ONE_VALID, temperature=1.0)[0]
|
||||||
|
hot = token_divergence(logits_of(s), logits_of(t), ONE_VALID, temperature=4.0)[0]
|
||||||
|
assert hot < cold
|
||||||
|
|
||||||
|
|
||||||
|
# ---- 3. §4.1 梯度爆炸演示 ----
|
||||||
|
|
||||||
|
|
||||||
|
def test_梯度爆炸_teacher概率趋0时student梯度暴涨():
|
||||||
|
# student 固定:对"采样 token"(id 0) 给最高 logit(模拟 on-policy 采到它)
|
||||||
|
base = [1.5, 0.5, 0.3, 0.2, 0.1]
|
||||||
|
epsilons = [1e-1, 1e-2, 1e-3, 1e-4, 1e-5, 1e-6]
|
||||||
|
grad_norms = []
|
||||||
|
for eps in epsilons:
|
||||||
|
student_logits = torch.tensor(base).reshape(1, 1, V).requires_grad_(True)
|
||||||
|
# teacher:token0 概率 = eps(越来越"厌恶"它),其余 (1-eps) 均分
|
||||||
|
t_probs = [(1 - eps) / (V - 1)] * V
|
||||||
|
t_probs[0] = eps
|
||||||
|
teacher_logits = torch.tensor(t_probs).log().reshape(1, 1, V)
|
||||||
|
loss, _ = token_divergence(student_logits, teacher_logits, ONE_VALID, beta=1.0)
|
||||||
|
loss.backward()
|
||||||
|
grad_norms.append(student_logits.grad.norm().item())
|
||||||
|
|
||||||
|
# 单调暴涨:teacher 越否定采样 token,student 梯度范数越大
|
||||||
|
for lo, hi in zip(grad_norms, grad_norms[1:]):
|
||||||
|
assert hi > lo
|
||||||
|
# 末端(π_T=1e-6)远超首端(π_T=1e-1)——§4.1 的可执行证据,
|
||||||
|
# 为层 5"有界乘子 π̂"的稳定性对照埋桩
|
||||||
|
assert grad_norms[-1] > 5 * grad_norms[0]
|
||||||
|
|
||||||
|
|
||||||
|
def test_全掩码batch显式报错():
|
||||||
|
all_masked = torch.full((1, 1), IGNORE_INDEX, dtype=torch.long)
|
||||||
|
with pytest.raises(ValueError, match="有效 completion"):
|
||||||
|
token_divergence(logits_of([0.2] * V), logits_of([0.2] * V), all_masked)
|
||||||
|
|
||||||
|
|
||||||
|
def test_beta越界报错():
|
||||||
|
with pytest.raises(ValueError, match="beta"):
|
||||||
|
token_divergence(
|
||||||
|
logits_of([0.2] * V), logits_of([0.2] * V), ONE_VALID, beta=2.0
|
||||||
|
)
|
||||||
Reference in New Issue
Block a user