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>
50 lines
1.7 KiB
Python
50 lines
1.7 KiB
Python
"""层 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.data import load_sft_dataset
|
|
|
|
MODEL_DIR = "/data/zym/outputs/sft_qwen3-0.6b_dapo1k" # 正式 1 epoch 的产物
|
|
|
|
# 不挂 teacher 解答(teacher_completions_path 缺省):只取题目做推理输入
|
|
ds = load_sft_dataset(
|
|
"data/dapo-math-17k-unique.parquet", subset_size=1000, seed=42
|
|
)
|
|
|
|
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("----- 生成(前 600 字符)-----")
|
|
print(completion[:600])
|
|
print()
|
|
|
|
print(
|
|
"判读:应为步骤化数学解答(markdown 风格、以 Answer: 行收尾的倾向);"
|
|
"乱码/复读/空输出 = 不通过。",
|
|
flush=True,
|
|
)
|