Files
ars-opd-rebuild/ars_opd/trainer.py
T
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

243 lines
12 KiB
Python
Raw 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.
"""训练编排(IO 边缘)。层 1 形态:掩码 SFT 损失 + 最小 HF Trainer 子类。
论文锚点:§3.1 式(1) 的标准交叉熵 SFT(监督目标是 teacher rollout,见 docs/02 §1)。
层 5 会在此模块长出式(8) 的 chunk 蒸馏损失与 KL 锚定;届时式(1) 路径保留为基线。
结构:损失算法收在纯张量函数 `sft_loss`(可用 toy 张量在 CPU 单测),
`SFTTrainer` 只做接线——模型前向、调 `sft_loss`、把指标并进日志,
其余一切(优化器、调度、DDP、checkpoint 保存)原样继承 HF Trainer。
"""
from __future__ import annotations
from typing import Any
import torch
import torch.nn.functional as F
from transformers import Trainer
from ars_opd.data import IGNORE_INDEX
# ---------------------------------------------------------------------------
# 损失算法(纯张量函数,对拍 distillation_trainer.py:751-759 + 2776-2874
# ---------------------------------------------------------------------------
def compute_prompt_length(attention_mask: torch.Tensor, labels: torch.Tensor) -> int:
"""batch 级 prompt 边界 = batch 内最短的"有效长度 - completion 长度"。
参数(B = batchT = padding 后长度):
- attention_mask: (B, T),左 padding 位置为 0
- labels: (B, T)completion 位置为 token id,其余为 -100
返回标量 pl。非显然约束:取 batch **最小值**是为了不切掉任何行的
completion token——左 padding 下所有序列右对齐,第 r 行的 completion 起点
索引是 T - comp_r ≥ prompt_r ≥ min,故切片 [pl:] 必然包含全部 completion
代价是长 prompt 行会漏进一些 prompt token,由 sft_loss 用 labels 重掩码兜住。
"""
full_lengths = attention_mask.sum(dim=1) # (B,) 每行非 padding 的 token 数
completion_lengths = (labels != IGNORE_INDEX).sum(dim=1) # (B,)
return int((full_lengths - completion_lengths).min().item())
def sft_loss(
logits: torch.Tensor,
input_ids: torch.Tensor,
labels: torch.Tensor,
attention_mask: torch.Tensor,
) -> tuple[torch.Tensor, int]:
"""式(1)completion 位置上的移位交叉熵。
参数(B = batchT = padding 后长度,V = 词表大小):
- logits: (B, T, V) 模型对全序列的输出
- input_ids / labels / attention_mask: (B, T)SFTCollator 的产物
返回 (标量 loss, 本 batch 有效 completion token 数)。
切片几何(docs/02 §2.4 的四行,此处为权威实现):
位置 t-1 的 logit 预测位置 t 的 token,故 logits 取 [pl-1, T-1) 、
targets 取 [pl, T),两段长度同为 T-pl,逐位对齐。
"""
pl = compute_prompt_length(attention_mask, labels)
if pl < 1:
# 需要位置 pl-1 的 logit 存在;pl=0 意味着某行完全没有 prompt token
# 数据管线出了问题(collator 保证 prompt 至少含模板 token
raise ValueError(f"prompt_length={pl} < 1,存在无 prompt 的数据行")
shifted_logits = logits[:, pl - 1 : -1, :] # (B, T, V) -> (B, T-pl, V)
targets = input_ids[:, pl:].clone() # (B, T-pl)clone: 下面要原地改写
# 重掩码:切片里漏进的 prompt token(长 prompt 行)与 padding 全部置 -100。
# labels 是权威掩码,切片几何只是省算力——正确性完全由这一步保证
invalid = labels[:, pl:] == IGNORE_INDEX # (B, T-pl)
targets[invalid] = IGNORE_INDEX
num_valid = int((~invalid).sum().item())
if num_valid == 0:
# 差异标注:参考实现(trainer:2810-2812)对非有限 loss 返回零梯度标量静默
# 继续;我们显式报错——collator 已挡下 prompt-only 行,走到这里仍全被
# 掩码只可能是数据坏了(如 teacher 返回空解答),必须暴露而非跳过
raise ValueError(
"本 batch 没有任何有效 completion token(全被 -100 掩码)。"
"检查 teacher 解答是否为空、completion 预算是否被截光。"
)
vocab = shifted_logits.shape[-1]
loss = F.cross_entropy(
shifted_logits.reshape(-1, vocab), # (B*(T-pl), V)
targets.reshape(-1), # (B*(T-pl),)
ignore_index=IGNORE_INDEX,
)
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_<t,x) ‖ π_T(·|y_<t,x))]。
本函数只管"给定两组**已对齐**的 logits,算散度标量"——on-policy 采样(y~π_θ)
与移位对齐([pl-1:-1])由调用方(DistillTrainer, U4)负责,复用 sft_loss 同一套
切片几何。故这里不吃 input_ids:散度是分布对分布,不需要目标 token,labels
仅用于定位有效位置。
参数(B=batch,T=已移位对齐长度,V=词表大小):
- student_logits: (B, T, V)student 前向输出(带梯度)
- teacher_logits: (B, T, V)teacher no_grad 前向输出(无梯度)
- labels: (B, T)completion 位为 token id、其余为 IGNORE_INDEX;只做有效位掩码
- beta: KL 方向(docs/03 §2.3 三副面孔)。0=前向 KL(π_T‖π_θ)、1=反向 KL(π_θ‖π_T)=
式(2)、(0,1)=JSD 插值
- temperature: softmax 前除进两侧 logits 的温度(§2.3),软化/锐化分布
返回 (per-token mean 散度标量, 有效 token 数)。
差异标注:参考实现(DT:2408-2491)含 top-k 稀疏 + 尾桶快路,我们只保留全词表这
一条精确路径(docs/03 §4 删除清单:本地同 tokenizer teacher 放得下)。因全词表
log_softmax 对有限 logits 恒有限,也不需要参考在 -inf 支持集上的 nan_to_num 兜底。
"""
if not 0.0 <= beta <= 1.0:
raise ValueError(f"beta 必须在 [0,1]0=前向/1=反向/中间=JSD),收到 {beta}")
# 温度除进 logits、softmax 之前(§2.3):调分布形状,不是等比缩概率
student_logits = student_logits / temperature
teacher_logits = teacher_logits / temperature
# 全程 log 域(§2.3 数值稳定性)。log_probs 两侧都要;probs 按分支只算用得上的
# 那一份——(B,T,V) 在真实规模下每份 ~2.5G(§5 显存账),默认 β=1 热路径只需
# student_probs,不materialize teacher_probs
student_log_probs = F.log_softmax(student_logits, dim=-1) # (B, T, V)
teacher_log_probs = F.log_softmax(teacher_logits, dim=-1) # (B, T, V)
if beta == 1.0:
# 反向 KL(π_θ‖π_T) = Σ_v π_θ (log π_θ log π_T) —— 式(2)
# 非显然约束(§4.1 梯度爆炸源):对 student logit 的梯度含 π_θ·log(π_θ/π_T),
# on-policy 采到 teacher 眼中烂 token(π_T→0)时 log 比值→∞,单 token 梯度可
# 炸掉整个 batch。这正是层 2 要亲眼观察、层 5 用有界乘子 π̂ 替换的病灶
per_token = (
student_log_probs.exp() * (student_log_probs - teacher_log_probs)
).sum(-1) # (B, T, V) -> (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.3reduction 实义为 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 接线
# ---------------------------------------------------------------------------
class SFTTrainer(Trainer):
"""最小 SFT Trainer:只重写 compute_loss 与 log,其余全部继承 HF Trainer。
用法(见 scripts/ 训练脚本):与 HF Trainer 完全同参构造,
data_collator 传 SFTCollatortrain_dataset 传 load_sft_dataset 的产物。
"""
def __init__(self, *args: Any, **kwargs: Any) -> None:
super().__init__(*args, **kwargs)
# 关 dropout(参考实现同款,docs/02 §3 保留清单):层 2+ 的蒸馏要求
# student/ref 两次前向可比,前向必须确定;Qwen3 默认无 dropout
# 此处是无害的统一前置
for module in self.model.modules():
if isinstance(module, torch.nn.Dropout):
module.p = 0.0
# 非显然约束:显式退出 HF 的新式梯度累积契约(v5 trainer.py:1977 文档
# 原话:"If you are not using num_items_in_batch ... overwrite
# self.model_accepts_loss_kwargs to False")。新式契约要求返回
# sum/全局token数,且依赖基类 compute_loss 尾部的 ×world_size 补偿
# trainer.py:2028)——我们整体重写了 compute_loss,那段补偿不会执行,
# 曾致 loss 与梯度 ÷4(2026-07-18,第二幕;第一幕是返回裸 mean 被 ×8,
# 全程记录见 docs/02 §5)。退出后回到经典契约:返回本微批 mean,
# Trainer 负责 ÷累积步数,日志跨卡平均,版本稳定
self.model_accepts_loss_kwargs = False
self._token_counts: list[int] = []
def compute_loss(
self,
model: Any,
inputs: dict[str, torch.Tensor],
return_outputs: bool = False,
num_items_in_batch: int | None = None,
) -> torch.Tensor | tuple[torch.Tensor, Any]:
# 非显然约束:不把 labels 传给模型前向。HF 模型收到 labels 会自己算
# "全序列移位 CE"并放进 outputs.loss,那会绕过 batch-min 切片与重掩码;
# 损失的权威实现只能有 sft_loss 一处
outputs = model(
input_ids=inputs["input_ids"],
attention_mask=inputs["attention_mask"],
)
loss, num_tokens = sft_loss(
outputs.logits,
inputs["input_ids"],
inputs["labels"],
inputs["attention_mask"],
)
# num_items_in_batch 有意忽略:已在 __init__ 退出新式契约(见彼处注释),
# 本函数返回微批 mean,÷累积步数由 Trainer.training_step 负责。
# 代价是微批按相同权重而非 token 数加权(百分之几的偏差,与参考实现同行为)
self._token_counts.append(num_tokens)
return (loss, outputs) if return_outputs else loss
def log(self, logs: dict[str, float], start_time: float | None = None) -> None:
"""在 HF 的常规日志里并入每步有效 token 数均值。
sanity run 时盯这个数:若它远小于预期(≈batch 内解答总长),说明掩码
把 completion 也吞了——loss 曲线看不出这种错,token 数看得出。
"""
if self._token_counts:
logs["sft/num_tokens_per_step"] = sum(self._token_counts) / len(
self._token_counts
)
self._token_counts = []
super().log(logs, start_time)