de7f36828a
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>
132 lines
5.2 KiB
Python
132 lines
5.2 KiB
Python
"""层 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
|
||
)
|