Files
ars-opd-rebuild/scripts/train_sft.py
T
iomgaa a0faec0df7 层1: 修复远程 sanity OOM——per_device batch 8→2、累积 2→8(全局 64 不变)
根因:150k 大词表下显存大头是 (B,T,V) logits 链(fp32 ~20G@B=8)与逐层激活,
均正比于 B 而与 0.6B 参数量无关。docs/02 §2.6 旧显存估算勘误入档;
train_sft.sh 加 expandable_segments 防碎片。

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 09:21:38 -04:00

149 lines
5.9 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=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)
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()