Files
iomgaa 0ca60ea93f 层2/U4: DistillTrainer 编排(on-policy 生成→teacher no_grad 前向→反向KL)
trainer.py(对应 docs/03 §5 U4):
- build_generated_batch: 纯函数,generate 输出重建 ids/attention/labels;
  "首个 eos 及之前有效"用 cumsum-self==0 实现,稳健对付 pad==eos / pad!=eos / 无eos
- DistillTrainer.compute_loss 四步:生成(no_grad,unwrap)→重建→双前向→移位+divergence
- 损失几何复用 compute_prompt_length(与 sft_loss 同款);梯度只经 student
- 构造时校验 teacher/student 同 tokenizer(白盒前提,比 DT:2876 更早)
- teacher eval+冻结、设备迁移推迟到 compute_loss;同层1退出梯度累积新式契约

test_distill.py(对应 docs/03 §2.5):
- build_generated_batch 6 例:eos居中屏蔽/无eos全监督/batch混长/pad==eos/立即eos/prompt段-100

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-19 05:16:57 -04:00

424 lines
21 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.
"""训练编排(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
def build_generated_batch(
prompts: torch.Tensor,
prompt_attention_mask: torch.Tensor,
gen_output: torch.Tensor,
eos_token_id: int,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""把 model.generate 的输出重建成 (input_ids, attention_mask, labels)。
对应 docs/03 §2.5"生成结果重建 input_ids/labels 写回"。on-policy 下 completion
不来自数据、而是 student 现场生成,故 labels 也在生成后现造:prompt 段全 -100、
生成段有效处填 token id,供 token_divergence 掩码。
参数(B=batchP=prompt padding 长度,G=本 batch 最大生成长度):
- prompts: (B, P) 左 padding 的 promptSFTCollator prompt_only 产物)
- prompt_attention_mask: (B, P) prompt 左 padding 位为 0
- gen_output: (B, P+G) model.generate 输出(前 P 列即 prompts,后 G 列是生成)
- eos_token_id: 生成终止符 id
返回 (input_ids (B,P+G), attention_mask (B,P+G), labels (B,P+G))。
非显然约束(生成段的右 padding 掩码):generate 对提前结束的序列在右侧补
padding 到 batch 最大长度。"首个 eos(含)之前有效"用 `cumsum - self == 0`
实现——它精确保留到首个 eos、屏蔽其后一切(无论其后是 eos 还是 pad,也无论
pad_token 是否等于 eos),避免"pad==eos 时把补位当解答"或"pad!=eos 时漏掉补位"
两种静默错误。无 eos(撞 max_new_tokens)则整段生成全有效。
"""
b, p = prompts.shape
gen_tokens = gen_output[:, p:] # (B, G) 纯生成段
is_eos = gen_tokens == eos_token_id # (B, G)
# cumsum - self:截至本位、其**之前**出现过的 eos 数;==0 即"首个 eos 及之前"
gen_valid = (is_eos.cumsum(dim=1) - is_eos.long()) == 0 # (B, G) bool
input_ids = gen_output
attention_mask = torch.cat(
[prompt_attention_mask, gen_valid.long()], dim=1
) # (B, P+G)
prompt_labels = torch.full(
(b, p), IGNORE_INDEX, dtype=torch.long, device=prompts.device
)
gen_labels = torch.where(
gen_valid, gen_tokens, torch.full_like(gen_tokens, IGNORE_INDEX)
) # 生成段:有效处填 token id,其余 -100
labels = torch.cat([prompt_labels, gen_labels], dim=1) # (B, P+G)
return input_ids, attention_mask, labels
# ---------------------------------------------------------------------------
# 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)
class DistillTrainer(Trainer):
"""层 2 white-box OPDon-policy 生成 → teacher no_grad 前向 → token 级反向 KL。
论文锚点:§3.1 式(2)。每个微批现场编排三步(docs/03 §2.2 主流程的精简版,
已按 §4 删除 buffer/稀疏路径/off-policy 抽签):
1. student 采样生成轨迹 y~π_θ(on-policyno_grad——只采样不回传);
2. student 带梯度前向 + teacher no_grad 前向,得两份全词表 logits;
3. token_divergence 算 KL,梯度只经 student 那一路。
与 SFTTrainer 的关系:损失几何(移位对齐、batch-min prompt_length、labels 重
掩码)完全同款,直接复用 compute_prompt_length;唯一差别是"completion 从哪来"
——SFT 读数据缓存,这里 student 现场生成。
构造(见 scripts/train_whitebox.py):像 HF Trainer 一样传 model(student)/args/
train_dataset/data_collator(prompt_only 的 SFTCollator),另用关键字传 teacher_model
与 teacher_tokenizer,以及蒸馏超参。
"""
def __init__(
self,
*args: Any,
teacher_model: Any,
teacher_tokenizer: Any,
beta: float = 1.0,
kl_temperature: float = 1.0,
gen_temperature: float = 1.0,
gen_top_p: float = 1.0,
max_new_tokens: int = 1024,
**kwargs: Any,
) -> None:
super().__init__(*args, **kwargs)
# student tokenizer 从 collator 取(prompt_only collator 必持有它)
student_tokenizer = self.data_collator.tokenizer
# 白盒前提校验(docs/03 §1、§2.6):KL 逐词表位对齐,同 tokenizer 才有意义。
# 构造时就炸——不像参考实现(DT:2876)拖到第一步前向才炸
if teacher_tokenizer.get_vocab() != student_tokenizer.get_vocab():
raise ValueError(
"teacher 与 student 的 tokenizer 词表不一致——白盒 KL 要求逐词表位"
"对应(docs/03 §1)。请换用与 student 同 tokenizer 的 teacher。"
)
self._student_tokenizer = student_tokenizer
# teachereval + 冻结参数 + 每进程一份副本(DDP 每卡一份,docs/03 §2.6)。
# 设备迁移推迟到 compute_loss——此刻 student 还没被 Trainer 放到卡上
self.teacher = teacher_model.eval()
for param in self.teacher.parameters():
param.requires_grad_(False)
self.beta = beta
self.kl_temperature = kl_temperature
self.gen_temperature = gen_temperature
self.gen_top_p = gen_top_p
self.max_new_tokens = max_new_tokens
# 关 dropout:生成、student 前向、teacher 前向三者须确定可比(同 SFTTrainer)
for module in self.model.modules():
if isinstance(module, torch.nn.Dropout):
module.p = 0.0
# 同层 1:显式退出 HF 新式梯度累积契约(docs/02 §5trainer.py:1977
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]:
prompts = inputs["prompts"]
prompt_attention_mask = inputs["prompt_attention_mask"]
# teacher 迁到 student 所在卡(一次性;.to 幂等,后续步是 no-op)
if self.teacher.device != prompts.device:
self.teacher = self.teacher.to(prompts.device)
# 1. on-policy 生成。no_gradGKD 标准做法是"采样一次、再 teacher-forcing
# 前向算分布",梯度经第 3 步的前向回传,不经采样本身。DDP 下须用 unwrap
# 后的模型(DDP 包装体不暴露 generate
unwrapped_model = self.accelerator.unwrap_model(model)
with torch.no_grad():
gen_output = unwrapped_model.generate(
input_ids=prompts,
attention_mask=prompt_attention_mask,
max_new_tokens=self.max_new_tokens,
do_sample=True,
temperature=self.gen_temperature,
top_p=self.gen_top_p,
pad_token_id=self._student_tokenizer.pad_token_id,
eos_token_id=self._student_tokenizer.eos_token_id,
)
# 2. 重建 input_ids/attention_mask/labelsprompt 段 -100、生成段有效处填 id)
input_ids, attention_mask, labels = build_generated_batch(
prompts,
prompt_attention_mask,
gen_output,
self._student_tokenizer.eos_token_id,
)
# 3. student 带梯度前向 + teacher no_grad 前向(两份全词表 logits)。
# 非显然约束:teacher no_grad 免掉的是它内部几十层激活的反向图(②),
# 但输出 logits(①)仍占满 (B,L,V) 显存——两份都要算进 §5 显存账
student_outputs = model(input_ids=input_ids, attention_mask=attention_mask)
with torch.no_grad():
teacher_logits = self.teacher(
input_ids=input_ids, attention_mask=attention_mask
).logits
# 4. 移位对齐(与 sft_loss 同款几何,docs/02 §2.4)→ token_divergence。
# 切片漏进的 prompt token 与生成段右 padding 由 labels 重掩码兜住
pl = compute_prompt_length(attention_mask, labels)
if pl < 1:
raise ValueError(f"prompt_length={pl} < 1,生成批次存在无 prompt 的行")
loss, num_tokens = token_divergence(
student_outputs.logits[:, pl - 1 : -1, :],
teacher_logits[:, pl - 1 : -1, :],
labels[:, pl:],
beta=self.beta,
temperature=self.kl_temperature,
)
self._token_counts.append(num_tokens)
return (loss, student_outputs) if return_outputs else loss
def log(self, logs: dict[str, float], start_time: float | None = None) -> None:
"""并入每步的生成 token 数均值:远程冒烟盯它,骤降=生成塌成空串。"""
if self._token_counts:
logs["distill/num_gen_tokens_per_step"] = sum(self._token_counts) / len(
self._token_counts
)
self._token_counts = []
super().log(logs, start_time)