0ca60ea93f
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>
424 lines
21 KiB
Python
424 lines
21 KiB
Python
"""训练编排(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 = batch,T = 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 = batch,T = 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.3:reduction 实义为 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=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 接线
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
class SFTTrainer(Trainer):
|
||
"""最小 SFT Trainer:只重写 compute_loss 与 log,其余全部继承 HF Trainer。
|
||
|
||
用法(见 scripts/ 训练脚本):与 HF Trainer 完全同参构造,
|
||
data_collator 传 SFTCollator,train_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 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)
|