"""层 2:white-box OPD 训练入口(由 train_whitebox.sh 经 torchrun 启动)。 自包含实验脚本:全部参数写死在下方 FULL 配置里,零参数复现;sanity 模式只是 对 FULL 的两处显式覆盖(50 步 + 独立输出目录)。 与层 1 train_sft.py 的结构差异:双模型(student + 本地 teacher)、prompt-only 数据(无 teacher 缓存,现场 on-policy 生成)、DistillTrainer 编排。 """ # ---- FSDP 前置块(必须在一切 transformers/accelerate import 之前,同 train_sft.py)---- 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 DistillConfig from ars_opd.data import SFTCollator, load_sft_dataset from ars_opd.trainer import DistillTrainer STUDENT_MODEL = "Qwen/Qwen3-0.6B" # 被训练的固定基线(同层 1,脚本级常量) FULL = DistillConfig( dataset_path="data/dapo-math-17k-unique.parquet", output_dir="/data/zym/outputs/whitebox_qwen3-0.6b_dapo1k", teacher_model="Qwen/Qwen3-4B", # 本地全词表 teacher(须与 student 同 tokenizer) subset_size=1000, seed=42, # 与层 1 一致:同一批题上对比 SFT 与蒸馏 max_prompt_length=1024, max_new_tokens=1024, # 与 max_prompt_length 之和 = 序列总长 T≈2048(§5 显存账) enable_thinking=False, beta=1.0, # 反向 KL = 式(2) kl_temperature=1.0, gen_temperature=1.0, # 纯采样自 π_θ(忠实 on-policy) gen_top_p=1.0, learning_rate=1e-6, # 论文 §5.1 蒸馏 lr;小步长也帮训练在梯度爆炸毛刺中存活 per_device_train_batch_size=4, # §5 估算,首次远程必须 nvidia-smi 核实不 OOM gradient_accumulation_steps=4, # 全局 batch = 4 × 4 卡 × 4 = 64(同层 1) num_train_epochs=1, max_steps=-1, bf16=True, logging_steps=1, save_steps=100, save_total_limit=2, report_to="none", ) def build_config() -> DistillConfig: """按命令行模式产出配置。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_prompt(dataset, collator, tokenizer) -> None: """训练前解码第一个 prompt 供肉眼核对(只在 rank0 打印一次)。 prompt-only 模式的自检重点:prompt 末尾应是生成引导符("...assistant\\n" + no-think 时的空 ),student 将从此续写。若末尾不对,生成的分布与 训练目标会错位。 """ batch = collator([dataset[0]]) prompt_ids = batch["prompts"][0] mask = batch["prompt_attention_mask"][0].bool() text = tokenizer.decode(prompt_ids[mask], skip_special_tokens=False) print( "=" * 30 + " 首个 prompt 自检(供 on-policy 生成)" + "=" * 30 + f"\n[{int(mask.sum())} tok,末尾应为生成引导符]\n…{text[-400:]}\n" + "=" * 80, flush=True, ) def main() -> None: cfg = build_config() rank0 = int(os.environ.get("RANK", "0")) == 0 # 加载顺序 fail-fast(同 train_sft.py):数据(毫秒级)→ tokenizer(几 MB)→ # 模型(GB 级)。层 2 无 teacher 缓存,数据是 prompt-only 子集 dataset = load_sft_dataset( cfg.dataset_path, cfg.dataset_split, cfg.subset_size, cfg.seed ) student_tokenizer = AutoTokenizer.from_pretrained(STUDENT_MODEL) teacher_tokenizer = AutoTokenizer.from_pretrained(cfg.teacher_model) collator = SFTCollator( student_tokenizer, max_prompt_length=cfg.max_prompt_length, enable_thinking=cfg.enable_thinking, prompt_only=True, # 层 2:只出 prompt 张量,completion 靠生成 ) if rank0: smoke_check_first_prompt(dataset, collator, student_tokenizer) # student fp32 + bf16 混合精度(同层 1);teacher 直接 bf16(只推理,省显存) student = AutoModelForCausalLM.from_pretrained(STUDENT_MODEL, dtype=torch.float32) teacher = AutoModelForCausalLM.from_pretrained( cfg.teacher_model, dtype=torch.bfloat16 ) args = TrainingArguments( output_dir=cfg.output_dir, remove_unused_columns=False, # 保住 messages 列供 collator(同层 1 注释) 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, ) trainer = DistillTrainer( model=student, args=args, train_dataset=dataset, data_collator=collator, teacher_model=teacher, teacher_tokenizer=teacher_tokenizer, # 构造时校验与 student 同词表 beta=cfg.beta, kl_temperature=cfg.kl_temperature, gen_temperature=cfg.gen_temperature, gen_top_p=cfg.gen_top_p, max_new_tokens=cfg.max_new_tokens, ) trainer.train() trainer.save_model() if rank0: student_tokenizer.save_pretrained(cfg.output_dir) print(f"训练完成,模型已存至 {cfg.output_dir}", flush=True) if __name__ == "__main__": main()