Files
iomgaa 404abc22bf 重构: load_sft_dataset 改吃散装参数(磨平接口回看记录的毛刺)
深模块修正:本函数只用 5 个字段,却索要整个 SFTConfig——层 1 无痛,但诊断脚本
被迫伪造 output_dir(4 处 /tmp/diag、outputs/_unused),层 2 更因 DistillConfig
无 teacher_completions_path 而无法复用。改收 dataset_path/split/subset_size/seed/
teacher_completions_path 五个散装参数(接口终于比实现轻)。

- data.py: 签名改散装参数;移除 TYPE_CHECKING 的 SFTConfig 依赖
- train_sft / diag_loss_probe / diag_collator: 仍持 SFTConfig(喂 collator),改调用点
- diag_generate / generate_teacher_completions: 只为 load 而造 config,直接丢弃、
  去掉伪造 output_dir,改传字面量
- 为 U5 层 2 训练脚本能直接 load_sft_dataset(distill_cfg 的字段) 铺路

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-19 05:21:16 -04:00

69 lines
2.5 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""损失探针:用未训练的预训练模型走完整管线,逐行算 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.dataset_path,
cfg.dataset_split,
cfg.subset_size,
cfg.seed,
cfg.teacher_completions_path,
)
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)