From 1976230250fd89dd569e3df461fb99ccf4391dbc Mon Sep 17 00:00:00 2001 From: iomgaa Date: Sat, 18 Jul 2026 10:14:31 -0400 Subject: [PATCH] =?UTF-8?q?=E5=B1=821:=20=E6=8D=9F=E5=A4=B1=E6=8E=A2?= =?UTF-8?q?=E9=92=88=E8=84=9A=E6=9C=AC=E2=80=94=E2=80=94=E9=A2=84=E8=AE=AD?= =?UTF-8?q?=E7=BB=83=E6=A8=A1=E5=9E=8B=E8=B5=B0=E7=AE=A1=E7=BA=BF=E9=80=90?= =?UTF-8?q?=E8=A1=8C=E7=AE=97=20CE=EF=BC=8C=E5=8C=BA=E5=88=86=E6=95=B0?= =?UTF-8?q?=E6=8D=AE=E9=97=AE=E9=A2=98=E4=B8=8E=E8=AE=AD=E7=BB=83=E7=8E=AF?= =?UTF-8?q?=E8=8A=82=E9=97=AE=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-Authored-By: Claude Fable 5 --- scripts/diag_loss_probe.py | 62 ++++++++++++++++++++++++++++++++++++++ 1 file changed, 62 insertions(+) create mode 100644 scripts/diag_loss_probe.py diff --git a/scripts/diag_loss_probe.py b/scripts/diag_loss_probe.py new file mode 100644 index 0000000..1094c25 --- /dev/null +++ b/scripts/diag_loss_probe.py @@ -0,0 +1,62 @@ +"""损失探针:用未训练的预训练模型走完整管线,逐行算 loss(层 1 疑点排查第二步)。 + +判读(训练日志初始 loss ≈ 7.5): +- 探针也 ≈ 7:管线一致,loss 高是数据/模型现实 → 去查数据(垃圾长文、乱码占比); +- 探针 ≈ 2-4:管线(本脚本与训练共用)没问题但训练环节另有妖 → 查训练循环差异。 +同时打印 HF 模型内建 CE(同一数学的独立实现)交叉验证 sft_loss。 + +远程运行(CPU 即可,约 1-2 分钟): + python -u scripts/diag_loss_probe.py +""" + +import torch +from transformers import AutoModelForCausalLM, AutoTokenizer + +from ars_opd.configs import SFTConfig +from ars_opd.data import SFTCollator, load_sft_dataset +from ars_opd.trainer import sft_loss + +MODEL = "Qwen/Qwen3-0.6B" + +cfg = SFTConfig( + dataset_path="data/dapo-math-17k-unique.parquet", + output_dir="/tmp/diag", + subset_size=8, + seed=42, + teacher_completions_path="data/teacher_completions_dapo1k_minimax-m3.jsonl", +) +ds = load_sft_dataset(cfg) +tok = AutoTokenizer.from_pretrained(MODEL) +collator = SFTCollator( + tok, + max_length=cfg.max_length, + max_prompt_length=cfg.max_prompt_length, + enable_thinking=False, +) +model = AutoModelForCausalLM.from_pretrained(MODEL, dtype=torch.float32) +model.eval() + +print(f"{'行':>3} {'sft_loss':>9} {'HF内建CE':>9} {'监督tok':>7} 解答开头") +total, total_n = 0.0, 0 +for i in range(len(ds)): + batch = collator([ds[i]]) + with torch.no_grad(): + out = model( + input_ids=batch["input_ids"], attention_mask=batch["attention_mask"] + ) + ours, n = sft_loss( + out.logits, batch["input_ids"], batch["labels"], batch["attention_mask"] + ) + # 交叉验证:HF 内建损失(labels 传入模型,内部自动移位)与 sft_loss + # 是同一数学的两个独立实现,单行 batch 下应当几乎相等 + hf = model( + input_ids=batch["input_ids"], + attention_mask=batch["attention_mask"], + labels=batch["labels"], + ).loss + head = ds[i]["messages"][-1]["content"][:40].replace("\n", " ") + print(f"{i:>3} {ours.item():>9.3f} {hf.item():>9.3f} {n:>7} {head}", flush=True) + total += ours.item() * n + total_n += n + +print(f"\n按 token 加权平均: {total / total_n:.3f}(对照训练日志初始 loss ≈ 7.5)", flush=True)