Files
iomgaa 404abc22bf 重构: load_sft_dataset 改吃散装参数(磨平接口回看记录的毛刺)
深模块修正:本函数只用 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>
2026-07-19 05:21:16 -04:00

155 lines
6.0 KiB
Python
Raw Permalink 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.
"""层 1SFT 基线训练入口(由 train_sft.sh 经 torchrun 启动,勿直接 python 运行)。
自包含实验脚本:全部参数写死在下方 FULL 配置里,零参数复现;sanity 模式只是
对 FULL 的两处显式覆盖(50 步 + 独立输出目录)。
"""
# ---- FSDP 前置块(必须在一切 transformers/accelerate import 之前)----
# 非显然约束:FSDP 的激活检查点开关是 accelerate 在 import 时读取的环境变量
# (参考实现 train_distillation.py:15-21 的著名坑);写在 import 后会静默无效。
# DDP 下本变量是无害 no-op——现在就位是为了未来换 4B 学生/FSDP 时只改此处一行,
# 且改完必须 nvidia-smi 实测显存验证生效(docs/02 §2.6)。
import os
os.environ.setdefault("FSDP_ACTIVATION_CHECKPOINTING", "false")
import dataclasses
import sys
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, TrainingArguments
from ars_opd.configs import SFTConfig
from ars_opd.data import IGNORE_INDEX, SFTCollator, load_sft_dataset
from ars_opd.trainer import SFTTrainer
STUDENT_MODEL = "Qwen/Qwen3-0.6B"
FULL = SFTConfig(
dataset_path="data/dapo-math-17k-unique.parquet",
output_dir="/data/zym/outputs/sft_qwen3-0.6b_dapo1k",
teacher_completions_path="data/teacher_completions_dapo1k_minimax-m3.jsonl",
subset_size=1000,
seed=42, # 非显然约束:与 generate_teacher_completions.py 一致,否则缓存大面积 miss
max_length=4096,
max_prompt_length=1024,
enable_thinking=False,
learning_rate=2e-5,
per_device_train_batch_size=2, # B=8 曾爆 80G:大头是 (B,T,V) logits 链与激活,见 SFTConfig 注释
gradient_accumulation_steps=8, # 全局 batch = 2 × 4 卡 × 8 = 64
num_train_epochs=1,
max_steps=-1,
lr_scheduler_type="linear",
warmup_ratio=0.0,
gradient_checkpointing=False,
bf16=True,
logging_steps=1,
save_steps=100,
save_total_limit=2,
report_to="none", # 层 1 先靠 tmux 实时日志;W&B 触发条件见 appendix C 表
)
def build_config() -> SFTConfig:
"""按命令行模式产出配置。frozen dataclass 的换参方式:replace 构造新实例。"""
mode = sys.argv[1] if len(sys.argv) > 1 else "full"
if mode == "full":
return FULL
if mode == "sanity":
return dataclasses.replace(
FULL, max_steps=50, output_dir=FULL.output_dir + "-sanity"
)
raise ValueError(f"未知模式 {mode!r},只接受 full / sanity")
def smoke_check_first_batch(dataset, collator, tokenizer) -> None:
"""训练前解码第一个 batch 供肉眼核对(只在 rank0 打印一次)。
单测用玩具 tokenizer 钉死了预算/边界的算法(tests/test_data.py),但真
tokenizer 的模板渲染只能在这里肉眼验证:掩码边界是否落在 assistant 起点、
no-think 时空 <think> 块是否在 prompt 侧。这是参考实现"一次性诊断打印"
的合理化版本(docs/02 §2.3)。
"""
batch = collator([dataset[0]])
ids, labels = batch["input_ids"][0], batch["labels"][0]
masked = labels == IGNORE_INDEX
prompt_text = tokenizer.decode(ids[masked], skip_special_tokens=False)
completion_text = tokenizer.decode(ids[~masked], skip_special_tokens=False)
print(
"=" * 30
+ " 首样本自检(人工核对掩码边界)"
+ "=" * 30
+ f"\n[prompt 段 | {int(masked.sum())} tok | 不产生 loss]\n"
+ f"…{prompt_text[-300:]}\n"
+ f"\n[completion 段 | {int((~masked).sum())} tok | 监督目标]\n"
+ f"{completion_text[:300]}\n"
+ "=" * 80,
flush=True,
)
def main() -> None:
cfg = build_config()
rank0 = int(os.environ.get("RANK", "0")) == 0
# 加载顺序刻意 fail-fast:数据(毫秒级,最易配错)→ tokenizer(几 MB)→
# 模型(GB 级下载)。teacher 缓存缺失要在下模型之前炸出来
dataset = load_sft_dataset(
cfg.dataset_path,
cfg.dataset_split,
cfg.subset_size,
cfg.seed,
cfg.teacher_completions_path,
)
tokenizer = AutoTokenizer.from_pretrained(STUDENT_MODEL)
collator = SFTCollator(
tokenizer,
max_length=cfg.max_length,
max_prompt_length=cfg.max_prompt_length,
enable_thinking=cfg.enable_thinking,
)
if rank0:
smoke_check_first_batch(dataset, collator, tokenizer)
model = AutoModelForCausalLM.from_pretrained(STUDENT_MODEL, dtype=torch.float32)
args = TrainingArguments(
output_dir=cfg.output_dir,
# 非显然约束:必须关掉列裁剪。HF Trainer 默认删除模型 forward 签名里
# 没有的数据列——"messages" 会被整列删光,collator 收到空字典且不报错
remove_unused_columns=False,
learning_rate=cfg.learning_rate,
per_device_train_batch_size=cfg.per_device_train_batch_size,
gradient_accumulation_steps=cfg.gradient_accumulation_steps,
num_train_epochs=cfg.num_train_epochs,
max_steps=cfg.max_steps,
lr_scheduler_type=cfg.lr_scheduler_type,
warmup_ratio=cfg.warmup_ratio,
gradient_checkpointing=cfg.gradient_checkpointing,
bf16=cfg.bf16,
seed=cfg.seed,
logging_steps=cfg.logging_steps,
logging_first_step=True,
save_strategy="steps",
save_steps=cfg.save_steps,
save_total_limit=cfg.save_total_limit,
report_to=cfg.report_to,
ddp_find_unused_parameters=False, # 全参训练无闲置参数,省一次全模型扫描
dataloader_num_workers=2, # collator 逐 batch 分词在 CPU,双 worker 与 GPU 重叠
)
trainer = SFTTrainer(
model=model,
args=args,
train_dataset=dataset,
data_collator=collator,
)
trainer.train()
trainer.save_model() # 终态模型(save_pretrained 格式,含 config
if rank0:
tokenizer.save_pretrained(cfg.output_dir)
print(f"训练完成,模型已存至 {cfg.output_dir}", flush=True)
if __name__ == "__main__":
main()