Files
ars-opd-rebuild/scripts/diag_loss_probe.py
T

63 lines
2.4 KiB
Python
Raw 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)
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)