"""层 1:SFT 基线训练入口(由 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 时空 块是否在 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) 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()