层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
+131
View File
@@ -0,0 +1,131 @@
"""层 2 / U2token 级散度单测(docs/03 §2、§4.1、§6.1)。
只测纯张量函数 token_divergenceDistillTrainer 类由远程冒烟验证。
三块:
1. 对拍 PyTorch 自带 KL(独立 oracle,非同式自证);
2. 反向/前向 KL 的方向性(mode-seeking vs mode-covering);
3. §4.1 梯度爆炸演示——teacher 概率趋 0 时 student 梯度暴涨(层 5 有界乘子的对照桩)。
构造技巧:softmax(log p) = pp 已归一),故用 `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) # 单个有效 tokenid 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
# 情形 Ateacher 尖、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
# 情形 Bstudent 尖、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)
# teachertoken0 概率 = 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 越否定采样 tokenstudent 梯度范数越大
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
)