Files
ars-opd-rebuild/scripts/train_whitebox.py
T
iomgaa b7f24d635e 层2/U5: white-box OPD 自包含训练脚本(train_whitebox.py + .sh)
对应 docs/03 §5 U5,对齐层 1 train_sft 骨架,换成层 2 装配:
- 双模型:student Qwen3-0.6B(fp32+bf16混训) + teacher Qwen3-4B(bf16 推理)
- prompt_only collator(无 teacher 缓存,现场 on-policy 生成)
- DistillTrainer 装配:teacher_model/teacher_tokenizer + beta/温度/生成参数
- FULL = §5 默认(beta=1、lr=1e-6、B=4×GA=4×4卡=全局64、max_new_tokens=1024)
- 首 prompt 自检(末尾须为生成引导符);sanity=50步冒烟
- .sh 前置清单强调白盒扛两份全词表 logits、OOM 阶梯、首跑下载 4B teacher

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-19 05:23:42 -04:00

156 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.
"""层 2white-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 时的空 <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()