层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>
This commit is contained in:
2026-07-19 05:23:42 -04:00
parent 404abc22bf
commit b7f24d635e
2 changed files with 186 additions and 0 deletions
+155
View File
@@ -0,0 +1,155 @@
"""层 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()
+31
View File
@@ -0,0 +1,31 @@
#!/usr/bin/env bash
# 层 2white-box OPD 训练(远程 gpu-a800-060 专用;本地不跑训练)。
#
# 用法(tmux 内执行,日志实时可查):
# bash scripts/train_whitebox.sh sanity # 50 步冒烟:看首 prompt 自检 + KL loss +
# # 生成 token 数;预期见 loss 毛刺(梯度爆炸实况)
# bash scripts/train_whitebox.sh # 正式:1k 子集 1 epoch
#
# 前置检查清单:
# 1. nvidia-smi 确认下方 GPUS 四张卡空闲(只许用 8 卡中的 4 张,严禁自动选卡)。
# ⚠️ 白盒显存比层 1 紧:student 训练全套 + teacher(4B) 推理副本 + **两份**全词表
# logits(student/teacher),§5 估算 B=4/T=2048 起步安全,但首跑必须盯 nvidia-smi;
# 若 OOM,降 per_device_train_batch_size 到 2,仍不够再开 gradient_checkpointing
# (改 DistillConfig,注意 checkpointing 与 generate 的 use_cache 交互)。
# 2. data/dapo-math-17k-unique.parquet 已在(层 2 无需 teacher 缓存,纯 prompt-only):
# scp data/dapo-math-17k-unique.parquet <远程>:/data/zym/ars-opd-rebuild/data/
# 3. 代码最新:git -C /data/zym/ars-opd-rebuild pull
# 4. 首跑会下载 teacher Qwen3-4BGB 级)到 HF_HOME,确保 /data 有空间
set -euo pipefail
cd "$(dirname "$0")/.." # 锚定仓库根
GPUS=0,1,2,3 # ⚠️ 改这里前先 nvidia-smi
MODE=${1:-full}
export CUDA_VISIBLE_DEVICES=$GPUS
export PYTHONUNBUFFERED=1 # 禁止日志缓存(CLAUDE.md §5
export PYTORCH_ALLOC_CONF=expandable_segments:True # 变长生成序列易碎片化,按需扩段
export HF_ENDPOINT=${HF_ENDPOINT:-https://hf-mirror.com}
export HF_HOME=${HF_HOME:-/data/zym/hf_cache} # 模型缓存落 /data,根分区已满
torchrun --nproc_per_node=4 --master_port=29572 scripts/train_whitebox.py "$MODE"