404abc22bf
深模块修正:本函数只用 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>
69 lines
2.5 KiB
Python
69 lines
2.5 KiB
Python
"""损失探针:用未训练的预训练模型走完整管线,逐行算 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)
|