Files
iomgaa de7f36828a 层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>
2026-07-19 04:00:25 -04:00

132 lines
5.2 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""层 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
)