Files
ars-opd-rebuild/scripts/train_sft.py
T
iomgaa c5a3b7d0bb 层1/T5: 自包含训练脚本 train_sft.sh + train_sft.py;层 1 代码收口
- train_sft.py: FSDP 环境变量前置块(import 前,DDP 下无害);fail-fast 加载
  顺序(数据→tokenizer→模型);首样本自检打印(真 tokenizer 掩码边界肉眼核对);
  remove_unused_columns=False 等非显然约束逐条注释;sanity 模式 = replace 覆盖
- train_sft.sh: 显式 CUDA_VISIBLE_DEVICES 4 卡、PYTHONUNBUFFERED、HF 镜像/缓存
  改道 /data,前置检查清单(含 scp 数据命令)
- 本地验证:fail-fast 到 teacher 缓存缺失处显式报错(941/1000,59 条真实命中
  反向证明 prompt_key 契约端到端成立)
- roadmap 存档点:层 1 代码完成,进入运行阶段

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 08:16:24 -04:00

149 lines
5.8 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.
"""层 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=8,
gradient_accumulation_steps=2, # 全局 batch = 8 × 4 卡 × 2 = 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)
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()