diff --git a/ars_opd/trainer.py b/ars_opd/trainer.py index e40fdda..b10f969 100644 --- a/ars_opd/trainer.py +++ b/ars_opd/trainer.py @@ -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=batch,P=prompt padding 长度,G=本 batch 最大生成长度): + - prompts: (B, P) 左 padding 的 prompt(SFTCollator 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 OPD:on-policy 生成 → teacher no_grad 前向 → token 级反向 KL。 + + 论文锚点:§3.1 式(2)。每个微批现场编排三步(docs/03 §2.2 主流程的精简版, + 已按 §4 删除 buffer/稀疏路径/off-policy 抽签): + 1. student 采样生成轨迹 y~π_θ(on-policy,no_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 + + # teacher:eval + 冻结参数 + 每进程一份副本(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 §5,trainer.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_grad:GKD 标准做法是"采样一次、再 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/labels(prompt 段 -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) diff --git a/tests/test_distill.py b/tests/test_distill.py new file mode 100644 index 0000000..980fbb6 --- /dev/null +++ b/tests/test_distill.py @@ -0,0 +1,91 @@ +"""层 2 / U4:生成重建纯逻辑单测(docs/03 §2.5)。 + +只测 build_generated_batch——把 model.generate 输出重建成 input_ids/attention/labels +的张量逻辑。DistillTrainer 的编排(生成→双前向→divergence)需真模型,由远程冒烟 +验证。重点盯"首个 eos 之后一律屏蔽"这条最易 off-by-one 的规则。 +""" + +import torch + +from ars_opd.data import IGNORE_INDEX +from ars_opd.trainer import build_generated_batch + +EOS = 99 +PAD = 0 + + +def test_eos在中间_其后全屏蔽(): + # 行内:prompt=[pad,u,u],生成=[a, EOS, pad];a 与 eos 有效,eos 后的 pad 无效 + prompts = torch.tensor([[PAD, 5, 5]]) + prompt_mask = torch.tensor([[0, 1, 1]]) + gen_output = torch.tensor([[PAD, 5, 5, 7, EOS, PAD]]) # (1, P+G)=(1,6) + + ids, attn, labels = build_generated_batch(prompts, prompt_mask, gen_output, EOS) + + assert ids.tolist() == gen_output.tolist() # input_ids 即生成全序列 + # prompt 段全 -100;生成段 [7, EOS, -100] + assert labels[0].tolist() == [IGNORE_INDEX] * 3 + [7, EOS, IGNORE_INDEX] + # attention:prompt 左 pad=0,生成段 eos 及之前=1、其后=0 + assert attn[0].tolist() == [0, 1, 1, 1, 1, 0] + + +def test_无eos撞max时整段生成有效(): + prompts = torch.tensor([[5, 5, 5]]) + prompt_mask = torch.tensor([[1, 1, 1]]) + gen_output = torch.tensor([[5, 5, 5, 8, 9, 10]]) # 生成三 token,无 eos + + _, attn, labels = build_generated_batch(prompts, prompt_mask, gen_output, EOS) + + assert labels[0].tolist() == [IGNORE_INDEX] * 3 + [8, 9, 10] # 全监督 + assert attn[0].tolist() == [1, 1, 1, 1, 1, 1] + + +def test_batch内不同生成长度_各自正确对齐(): + # 行0 提前 eos(右侧被补 pad 到 batch 宽度);行1 撞 max。二者共用同一 (B,P+G) + prompts = torch.tensor([[PAD, 5, 5], [5, 5, 5]]) + prompt_mask = torch.tensor([[0, 1, 1], [1, 1, 1]]) + gen_output = torch.tensor( + [ + [PAD, 5, 5, 7, EOS, PAD], # 行0:生成 [7, EOS],末位 pad 补齐 + [5, 5, 5, 8, 9, 10], # 行1:生成 [8, 9, 10] + ] + ) + + _, attn, labels = build_generated_batch(prompts, prompt_mask, gen_output, EOS) + + assert labels[0].tolist() == [IGNORE_INDEX] * 3 + [7, EOS, IGNORE_INDEX] + assert labels[1].tolist() == [IGNORE_INDEX] * 3 + [8, 9, 10] + assert attn[0].tolist() == [0, 1, 1, 1, 1, 0] + assert attn[1].tolist() == [1, 1, 1, 1, 1, 1] + + +def test_pad等于eos也不误判(): + # 关键 corner:pad_token == eos_token。首个 eos 有效、其后补位的 eos 全屏蔽 + prompts = torch.tensor([[5, 5, 5]]) + prompt_mask = torch.tensor([[1, 1, 1]]) + gen_output = torch.tensor([[5, 5, 5, 7, EOS, EOS]]) # 末位补的 pad 恰好==eos + + _, attn, labels = build_generated_batch(prompts, prompt_mask, gen_output, EOS) + + assert labels[0].tolist() == [IGNORE_INDEX] * 3 + [7, EOS, IGNORE_INDEX] + assert attn[0].tolist() == [1, 1, 1, 1, 1, 0] # 第二个 eos 被当补位屏蔽 + + +def test_立即eos_只留一个token(): + prompts = torch.tensor([[5, 5, 5]]) + prompt_mask = torch.tensor([[1, 1, 1]]) + gen_output = torch.tensor([[5, 5, 5, EOS, PAD, PAD]]) # 第一步就 eos + + _, attn, labels = build_generated_batch(prompts, prompt_mask, gen_output, EOS) + + assert labels[0].tolist() == [IGNORE_INDEX] * 3 + [EOS, IGNORE_INDEX, IGNORE_INDEX] + assert attn[0].tolist() == [1, 1, 1, 1, 0, 0] + + +def test_prompt段恒为负100(): + prompts = torch.tensor([[PAD, PAD, 5, 5]]) + prompt_mask = torch.tensor([[0, 0, 1, 1]]) + gen_output = torch.tensor([[PAD, PAD, 5, 5, 8, 9]]) + + _, _, labels = build_generated_batch(prompts, prompt_mask, gen_output, EOS) + assert labels[0, :4].tolist() == [IGNORE_INDEX] * 4 # prompt 段(含左 pad)全 -100