层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:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user