层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>
This commit is contained in:
2026-07-19 05:16:57 -04:00
parent e5a28e8e77
commit 0ca60ea93f
2 changed files with 272 additions and 0 deletions
+181
View File
@@ -171,6 +171,52 @@ def token_divergence(
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 接线
# ---------------------------------------------------------------------------
@@ -240,3 +286,138 @@ class SFTTrainer(Trainer):
)
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)