8b362eae09
首冒烟发现:sanity 的 loss 平滑、无预期毛刺,因 HF 默认 max_grad_norm=1.0 把 反向 KL 的梯度爆炸(§4.1,实测 grad_norm 14→2 是裁剪前范数)默默压平了——正是 本项目要堵的"静默行为"。 - configs.py: DistillConfig 加 max_grad_norm=1.0(默认=原 HF 行为),docstring 讲清 它是 §4.1 爆炸的隐形稳定器、日志 grad_norm 是裁剪前值;__post_init__ 校验 >0 - train_whitebox.py: FULL 显式写出、TrainingArguments 传入;build_config 加 noclip 模式(max_grad_norm=1e9≈关裁剪 + lr 5× + 15 步)暴露原始爆炸供教学对照 - .sh: 用法加 noclip 模式说明 Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
169 lines
6.6 KiB
Python
169 lines
6.6 KiB
Python
"""层 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,
|
||
max_grad_norm=1.0, # 显式写出这个此前静默的稳定器(§4.1 爆炸靠它压平,见 config 注释)
|
||
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"
|
||
)
|
||
if mode == "noclip":
|
||
# §4.1 教学对照:关闭裁剪 + 稍抬 lr,暴露反向 KL 原始爆炸。max_grad_norm
|
||
# 设远高于实测范数(~14)故永不触发≈无裁剪;lr 5×放大让爆炸在 loss 上可见。
|
||
# 与 sanity(裁到 1.0、lr 1e-6 的平滑曲线)并排 = 白盒脆弱性活教材,层 5 对照
|
||
return dataclasses.replace(
|
||
FULL,
|
||
max_steps=15,
|
||
max_grad_norm=1e9,
|
||
learning_rate=5e-6,
|
||
output_dir=FULL.output_dir + "-noclip",
|
||
)
|
||
raise ValueError(f"未知模式 {mode!r},只接受 full / sanity / noclip")
|
||
|
||
|
||
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,
|
||
max_grad_norm=cfg.max_grad_norm,
|
||
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()
|