From e472a66959f5c54aaf83ccc32f5f4229d2f510d2 Mon Sep 17 00:00:00 2001 From: iomgaa Date: Sat, 18 Jul 2026 10:37:13 -0400 Subject: [PATCH] =?UTF-8?q?=E5=B1=821:=20checkpoint=20=E7=94=9F=E6=88=90?= =?UTF-8?q?=E6=B5=8B=E8=AF=95=E8=84=9A=E6=9C=AC=EF=BC=88=E5=85=B3=E8=B4=A6?= =?UTF-8?q?=E5=88=A4=E6=8D=AE=203=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-Authored-By: Claude Fable 5 --- scripts/diag_generate.py | 49 ++++++++++++++++++++++++++++++++++++++++ 1 file changed, 49 insertions(+) create mode 100644 scripts/diag_generate.py diff --git a/scripts/diag_generate.py b/scripts/diag_generate.py new file mode 100644 index 0000000..cc96203 --- /dev/null +++ b/scripts/diag_generate.py @@ -0,0 +1,49 @@ +"""层 1 关账判据 3:训练后 checkpoint 能被 from_pretrained 加载并生成通顺解答。 + +远程运行(CPU 即可,0.6B 生成 512 token 约 1-2 分钟): + python -u scripts/diag_generate.py +""" + +import torch +from transformers import AutoModelForCausalLM, AutoTokenizer + +from ars_opd.configs import SFTConfig +from ars_opd.data import load_sft_dataset + +MODEL_DIR = "/data/zym/outputs/sft_qwen3-0.6b_dapo1k" # 正式 1 epoch 的产物 + +cfg = SFTConfig( + dataset_path="data/dapo-math-17k-unique.parquet", + output_dir="/tmp/diag", + subset_size=1000, + seed=42, + # 不挂 teacher 解答:只取题目做推理输入 +) +ds = load_sft_dataset(cfg) + +tok = AutoTokenizer.from_pretrained(MODEL_DIR) +model = AutoModelForCausalLM.from_pretrained(MODEL_DIR, dtype=torch.float32) +model.eval() + +# 取子集第 900+ 行附近的题(训练时见过,此处只验"会不会说话"不验泛化) +for i in (900, 950): + prompt = tok.apply_chat_template( + ds[i]["messages"], + tokenize=False, + add_generation_prompt=True, + enable_thinking=False, # 必须与训练取值一致(docs/02 §2.3 边界契约) + ) + inputs = tok(prompt, return_tensors="pt", add_special_tokens=False) + with torch.no_grad(): + out = model.generate( + **inputs, max_new_tokens=512, do_sample=False, temperature=None, top_p=None + ) + completion = tok.decode(out[0][inputs["input_ids"].shape[1] :], skip_special_tokens=True) + print(f"===== 样本 {i} 题目 =====") + print(ds[i]["messages"][-1]["content"][120:280], "…") + print(f"----- 生成(前 600 字符)-----") + print(completion[:600]) + print() + +print("判读:应为步骤化数学解答(markdown 风格、以 Answer: 行收尾的倾向);" + "乱码/复读/空输出 = 不通过。", flush=True)