From b7f24d635edb764396feed623a366602e4836201 Mon Sep 17 00:00:00 2001 From: iomgaa Date: Sun, 19 Jul 2026 05:23:42 -0400 Subject: [PATCH] =?UTF-8?q?=E5=B1=822/U5:=20white-box=20OPD=20=E8=87=AA?= =?UTF-8?q?=E5=8C=85=E5=90=AB=E8=AE=AD=E7=BB=83=E8=84=9A=E6=9C=AC=EF=BC=88?= =?UTF-8?q?train=5Fwhitebox.py=20+=20.sh=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 对应 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) --- scripts/train_whitebox.py | 155 ++++++++++++++++++++++++++++++++++++++ scripts/train_whitebox.sh | 31 ++++++++ 2 files changed, 186 insertions(+) create mode 100644 scripts/train_whitebox.py create mode 100755 scripts/train_whitebox.sh diff --git a/scripts/train_whitebox.py b/scripts/train_whitebox.py new file mode 100644 index 0000000..a07a6e3 --- /dev/null +++ b/scripts/train_whitebox.py @@ -0,0 +1,155 @@ +"""层 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() diff --git a/scripts/train_whitebox.sh b/scripts/train_whitebox.sh new file mode 100755 index 0000000..6b3f528 --- /dev/null +++ b/scripts/train_whitebox.sh @@ -0,0 +1,31 @@ +#!/usr/bin/env bash +# 层 2:white-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-4B(GB 级)到 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"