"""层 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 )