From ee16e298469b719223a30f5fb4190666d834a8f3 Mon Sep 17 00:00:00 2001 From: iomgaa Date: Sat, 18 Jul 2026 10:18:09 -0400 Subject: [PATCH] =?UTF-8?q?=E5=B1=821:=20=E4=BF=AE=E5=A4=8D=20HF=20?= =?UTF-8?q?=E6=A2=AF=E5=BA=A6=E7=B4=AF=E7=A7=AF=E5=A5=91=E7=BA=A6=E5=9D=91?= =?UTF-8?q?=E2=80=94=E2=80=94compute=5Floss=20=E6=8C=89=E6=96=B0=E5=BC=8F?= =?UTF-8?q?=E5=A5=91=E7=BA=A6=E8=BF=94=E5=9B=9E=20sum/num=5Fitems=5Fin=5Fb?= =?UTF-8?q?atch?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 根因(探针定案):预训练模型真实 CE≈0.85,训练日志 7.5≈0.94×8(累积步数)。 Qwen3 forward 接受 loss_kwargs → Trainer 走新式契约不再除以累积步数,我们 返回裸 mean 导致日志与梯度同放大 8 倍。sft_loss 纯函数不动,适配收口在 compute_loss;docs/02 §5 勘误起点预期(~0.85)并记录此坑。 Co-Authored-By: Claude Fable 5 --- ars_opd/trainer.py | 8 ++++++++ docs/02-sft-baseline.md | 2 +- 2 files changed, 9 insertions(+), 1 deletion(-) diff --git a/ars_opd/trainer.py b/ars_opd/trainer.py index d209156..e467a7d 100644 --- a/ars_opd/trainer.py +++ b/ars_opd/trainer.py @@ -133,6 +133,14 @@ class SFTTrainer(Trainer): inputs["labels"], inputs["attention_mask"], ) + # 非显然约束:HF 梯度累积契约(transformers 4.46 修正后的新式)。模型 + # forward 接受 loss_kwargs(Qwen3 是)时,Trainer 不再帮你除以累积步数, + # 而是期望 compute_loss 返回 sum/num_items_in_batch(整个累积组的总有效 + # token 数),各微批直接相加得到精确的全局 per-token 均值。返回裸 mean + # 会导致日志与梯度同时放大"累积步数"倍——2026-07-18 远程 sanity 的 + # loss 7.5 ≈ 0.94×8 正是此坑(诊断记录见 docs/02 §5) + if num_items_in_batch is not None: + loss = loss * num_tokens / num_items_in_batch # mean → sum/N_total self._token_counts.append(num_tokens) return (loss, outputs) if return_outputs else loss diff --git a/docs/02-sft-baseline.md b/docs/02-sft-baseline.md index a030cf6..0cabec5 100644 --- a/docs/02-sft-baseline.md +++ b/docs/02-sft-baseline.md @@ -83,5 +83,5 @@ F.cross_entropy(..., ignore_index=-100) ## 5. 验证方式 1. **本地(CPU)**:collator 单测对拍参考行为——构造超长解答样本断言 prompt 未被截空;断言 -100 位置分布;断言 enable_thinking 两种取值下边界正确。数据加载单测:断言 `except:pass` 已变显式报错。 -2. **远程**:先 sanity(官方 SFTTrainer 或我们管线跑 50 步)确认 loss 从 ~2-3 稳定下降;再全量 1 epoch。W&B 或实时日志盯 `sft/loss` 与有效 token 数。 +2. **远程**:先 sanity(官方 SFTTrainer 或我们管线跑 50 步)确认 loss 稳定下降;再全量 1 epoch。W&B 或实时日志盯 `sft/loss` 与有效 token 数。⚠️ 两条实测勘误(2026-07-18):① 起点不是想象的 ~2-3——预训练 Qwen3-0.6B 对 M3 风格数学文本的真实 CE ≈ **0.85**(探针 scripts/diag_loss_probe.py 实测),健康曲线 ≈ 0.9→0.4;② **HF 梯度累积契约坑**:模型 forward 接受 loss_kwargs 时(Qwen3 是),自定义 compute_loss 必须返回 `sum/num_items_in_batch` 而非裸 mean,否则日志与梯度都放大"累积步数"倍(首跑 loss 7.5 ≈ 0.94×8 即此坑;修复在 trainer.py compute_loss)。 3. 关账判据:1k 子集 1 epoch 跑完,loss 曲线正常,checkpoint 能被 `from_pretrained` 加载并生成通顺文本。