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

120 lines
4.7 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.
"""诊断脚本:逐环检验 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())