From de7f36828a93b84c45ecb7d59d051875f96c54d8 Mon Sep 17 00:00:00 2001 From: iomgaa Date: Sun, 19 Jul 2026 04:00:25 -0400 Subject: [PATCH] =?UTF-8?q?=E5=B1=822/U2:=20=E7=BA=AF=E9=80=BB=E8=BE=91=20?= =?UTF-8?q?token=5Fdivergence=EF=BC=88=E5=BC=8F(2)=20=E5=85=A8=E8=AF=8D?= =?UTF-8?q?=E8=A1=A8=20KL/JSD=EF=BC=89+=20=E6=A2=AF=E5=BA=A6=E7=88=86?= =?UTF-8?q?=E7=82=B8=E6=BC=94=E7=A4=BA=E5=8D=95=E6=B5=8B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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) --- ars_opd/trainer.py | 80 ++++++++++++++++++++++++ tests/test_divergence.py | 131 +++++++++++++++++++++++++++++++++++++++ 2 files changed, 211 insertions(+) create mode 100644 tests/test_divergence.py diff --git a/ars_opd/trainer.py b/ars_opd/trainer.py index 0f32363..e40fdda 100644 --- a/ars_opd/trainer.py +++ b/ars_opd/trainer.py @@ -91,6 +91,86 @@ def sft_loss( 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_ (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 接线 # --------------------------------------------------------------------------- diff --git a/tests/test_divergence.py b/tests/test_divergence.py new file mode 100644 index 0000000..10fca86 --- /dev/null +++ b/tests/test_divergence.py @@ -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 + )