层1: collator 对齐诊断脚本——排查远程 loss 7.5 与 \50118 解码异常(逐环检验缓存/模板/分词/切片)

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
2026-07-18 10:05:46 -04:00
parent a0faec0df7
commit ff572bf4b9
+119
View File
@@ -0,0 +1,119 @@
"""诊断脚本:逐环检验 collator 对齐链(层 1 关账前的疑点排查)。
背景:远程 sanity 中初始 loss ~7.5(预期 ~2-3),且首样本自检的 completion 段
出现 `$k \\50118$`teacher 原文是 `$k \\leq 2018$`)。本脚本把
"缓存文本 → 模板渲染 → 分词 → 边界切片 → 解码"逐环单测,定位腐坏点。
远程运行(CPU 即可,tokenizer 用已有 HF 缓存):
python -u scripts/diag_collator.py
"""
import hashlib
import sys
from transformers import AutoTokenizer
from ars_opd.configs import SFTConfig
from ars_opd.data import IGNORE_INDEX, SFTCollator, load_sft_dataset
CACHE = "data/teacher_completions_dapo1k_minimax-m3.jsonl"
DATASET = "data/dapo-math-17k-unique.parquet"
MODEL = "Qwen/Qwen3-0.6B"
def check(name: str, ok: bool, detail: str = "") -> bool:
print(f"[{'通过' if ok else '失败'}] {name}" + (f" —— {detail}" if detail else ""), flush=True)
return ok
def first_diff(a: str, b: str) -> int:
n = min(len(a), len(b))
for i in range(n):
if a[i] != b[i]:
return i
return -1 if len(a) == len(b) else n
def main() -> None:
# 环 0:缓存文件指纹(与本地对比,排除 scp 传坏/版本不一致)
digest = hashlib.sha256(open(CACHE, "rb").read()).hexdigest()
print(f"缓存文件 sha256: {digest[:16]}… (与本地对比)", flush=True)
cfg = SFTConfig(
dataset_path=DATASET,
output_dir="/tmp/diag",
subset_size=5,
seed=42,
teacher_completions_path=CACHE,
)
ds = load_sft_dataset(cfg)
msgs = ds[0]["messages"]
completion_text = msgs[-1]["content"]
# 环 1:本机数据管线出来的 teacher 文本是否干净
check(
"环1 缓存→数据集文本干净",
"\\leq 2018" in completion_text and "\\50118" not in completion_text,
f"开头: {completion_text[:60]!r}",
)
tok = AutoTokenizer.from_pretrained(MODEL)
fp = tok.apply_chat_template(
msgs[:-1], tokenize=False, add_generation_prompt=True, enable_thinking=False
)
ff = tok.apply_chat_template(
msgs, tokenize=False, add_generation_prompt=False, enable_thinking=False
)
# 环 2:完整渲染必须以 prompt 渲染为前缀(collator 边界法的前提假设!)
prefix_ok = ff.startswith(fp)
check("环2 完整渲染以 prompt 渲染为前缀", prefix_ok)
if not prefix_ok:
i = first_diff(ff, fp)
print(f" 首个分歧在第 {i} 字符:\n"
f" prompt 渲染: …{fp[max(0, i - 60) : i + 60]!r}\n"
f" 完整渲染: …{ff[max(0, i - 60) : i + 60]!r}", flush=True)
# 环 3:完整渲染中 teacher 文本是否原样存在(模板会不会改写 content)
check(
"环3 完整渲染保留 teacher 原文",
"\\leq 2018" in ff and "\\50118" not in ff,
"" if "\\leq 2018" in ff else "模板改写了 assistant content",
)
# 环 4:token 级前缀(坑一:拼接稳定性)
full_ids = tok(ff, add_special_tokens=False)["input_ids"]
fp_ids = tok(fp, add_special_tokens=False)["input_ids"]
tok_prefix_ok = full_ids[: len(fp_ids)] == fp_ids
check("环4 token 级前缀一致(无跨界合并)", tok_prefix_ok)
if not tok_prefix_ok:
i = next(k for k in range(len(fp_ids)) if full_ids[k] != fp_ids[k])
lo, hi = max(0, i - 3), i + 4
print(f" 首个分歧在 token {i}/{len(fp_ids)}\n"
f" prompt 侧: {[tok.decode([t]) for t in fp_ids[lo:hi]]}\n"
f" 完整侧: {[tok.decode([t]) for t in full_ids[lo:hi]]}", flush=True)
# 环 5collator 全流程后,completion 解码应等于完整渲染去掉 prompt 的尾段前缀
collator = SFTCollator(
tok,
max_length=cfg.max_length,
max_prompt_length=cfg.max_prompt_length,
enable_thinking=False,
)
batch = collator([ds[0]])
ids, labels = batch["input_ids"][0], batch["labels"][0]
comp_decoded = tok.decode(ids[labels != IGNORE_INDEX], skip_special_tokens=False)
expected_tail = ff[len(fp) :] if prefix_ok else "(环2 已失败,无期望值)"
tail_ok = prefix_ok and expected_tail.startswith(comp_decoded[:200])
check("环5 completion 解码 == 渲染尾段", tail_ok)
if prefix_ok and not tail_ok:
i = first_diff(comp_decoded, expected_tail)
print(f" 首个分歧在第 {i} 字符:\n"
f" 解码: …{comp_decoded[max(0, i - 50) : i + 50]!r}\n"
f" 期望: …{expected_tail[max(0, i - 50) : i + 50]!r}", flush=True)
print("\n诊断完成。把全部输出贴回对话。", flush=True)
if __name__ == "__main__":
sys.exit(main())