Compare commits
2 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| b7f24d635e | |||
| 404abc22bf |
+23
-13
@@ -18,14 +18,11 @@ import ast
|
|||||||
import hashlib
|
import hashlib
|
||||||
import json
|
import json
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING, Any
|
from typing import Any
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from datasets import Dataset, load_dataset
|
from datasets import Dataset, load_dataset
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from ars_opd.configs import SFTConfig
|
|
||||||
|
|
||||||
# F.cross_entropy 的 ignore_index 默认值;标了它的位置不产生 loss
|
# F.cross_entropy 的 ignore_index 默认值;标了它的位置不产生 loss
|
||||||
IGNORE_INDEX = -100
|
IGNORE_INDEX = -100
|
||||||
|
|
||||||
@@ -165,23 +162,36 @@ def attach_teacher_completions(dataset: Dataset, jsonl_path: str) -> Dataset:
|
|||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
def load_sft_dataset(cfg: "SFTConfig") -> Dataset:
|
def load_sft_dataset(
|
||||||
"""层 1 数据管线入口:加载 → 归一 → 抽子集 → 挂 teacher 解答。
|
dataset_path: str,
|
||||||
|
dataset_split: str = "train",
|
||||||
|
subset_size: int | None = None,
|
||||||
|
seed: int = 42,
|
||||||
|
teacher_completions_path: str | None = None,
|
||||||
|
) -> Dataset:
|
||||||
|
"""数据管线入口:加载 → 归一 → 抽子集 →(可选)挂 teacher 解答。
|
||||||
|
|
||||||
返回只含 ``messages`` 一列的 Dataset,每行末轮是 assistant(可直接喂 SFTCollator)。
|
收散装参数而非整个 config(深模块:本函数只用这 5 个字段,不该索要一整个
|
||||||
|
SFTConfig)。这样层 1(SFTConfig)、层 2(DistillConfig,无 teacher 缓存)、
|
||||||
|
诊断脚本都能直接调,无需伪造无关字段。teacher_completions_path=None 时
|
||||||
|
返回 prompt-only 数据集(末轮 user,供 on-policy 生成);给了则挂 teacher
|
||||||
|
解答(末轮 assistant,供 SFT)。
|
||||||
|
|
||||||
|
返回只含 ``messages`` 一列的 Dataset。
|
||||||
"""
|
"""
|
||||||
ds = _load_raw(cfg.dataset_path, cfg.dataset_split)
|
ds = _load_raw(dataset_path, dataset_split)
|
||||||
ds = ds.map(
|
ds = ds.map(
|
||||||
to_messages,
|
to_messages,
|
||||||
remove_columns=[c for c in ds.column_names if c != "messages"],
|
remove_columns=[c for c in ds.column_names if c != "messages"],
|
||||||
)
|
)
|
||||||
if cfg.subset_size is not None and cfg.subset_size < len(ds):
|
if subset_size is not None and subset_size < len(ds):
|
||||||
# 非显然约束:抽子集必须在挂接 teacher 解答之前、且由 seed 完全确定——
|
# 非显然约束:抽子集必须在挂接 teacher 解答之前、且由 seed 完全确定——
|
||||||
# teacher.py 生成缓存时会走完全相同的"加载→归一→抽子集"路径,两侧 seed
|
# teacher.py 生成缓存时会走完全相同的"加载→归一→抽子集"路径,两侧 seed
|
||||||
# 一致才能得到同一批题;否则 attach 处大面积缓存 miss 报错。
|
# 一致才能得到同一批题;否则 attach 处大面积缓存 miss 报错。层 2 与层 1
|
||||||
ds = ds.shuffle(seed=cfg.seed).select(range(cfg.subset_size))
|
# 用同 seed 同 subset_size,才能在同一批题上对比 SFT 与蒸馏。
|
||||||
if cfg.teacher_completions_path is not None:
|
ds = ds.shuffle(seed=seed).select(range(subset_size))
|
||||||
ds = attach_teacher_completions(ds, cfg.teacher_completions_path)
|
if teacher_completions_path is not None:
|
||||||
|
ds = attach_teacher_completions(ds, teacher_completions_path)
|
||||||
return ds
|
return ds
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -46,7 +46,13 @@ def main() -> None:
|
|||||||
seed=42,
|
seed=42,
|
||||||
teacher_completions_path=CACHE,
|
teacher_completions_path=CACHE,
|
||||||
)
|
)
|
||||||
ds = load_sft_dataset(cfg)
|
ds = load_sft_dataset(
|
||||||
|
cfg.dataset_path,
|
||||||
|
cfg.dataset_split,
|
||||||
|
cfg.subset_size,
|
||||||
|
cfg.seed,
|
||||||
|
cfg.teacher_completions_path,
|
||||||
|
)
|
||||||
msgs = ds[0]["messages"]
|
msgs = ds[0]["messages"]
|
||||||
completion_text = msgs[-1]["content"]
|
completion_text = msgs[-1]["content"]
|
||||||
|
|
||||||
|
|||||||
@@ -7,19 +7,14 @@
|
|||||||
import torch
|
import torch
|
||||||
from transformers import AutoModelForCausalLM, AutoTokenizer
|
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||||
|
|
||||||
from ars_opd.configs import SFTConfig
|
|
||||||
from ars_opd.data import load_sft_dataset
|
from ars_opd.data import load_sft_dataset
|
||||||
|
|
||||||
MODEL_DIR = "/data/zym/outputs/sft_qwen3-0.6b_dapo1k" # 正式 1 epoch 的产物
|
MODEL_DIR = "/data/zym/outputs/sft_qwen3-0.6b_dapo1k" # 正式 1 epoch 的产物
|
||||||
|
|
||||||
cfg = SFTConfig(
|
# 不挂 teacher 解答(teacher_completions_path 缺省):只取题目做推理输入
|
||||||
dataset_path="data/dapo-math-17k-unique.parquet",
|
ds = load_sft_dataset(
|
||||||
output_dir="/tmp/diag",
|
"data/dapo-math-17k-unique.parquet", subset_size=1000, seed=42
|
||||||
subset_size=1000,
|
|
||||||
seed=42,
|
|
||||||
# 不挂 teacher 解答:只取题目做推理输入
|
|
||||||
)
|
)
|
||||||
ds = load_sft_dataset(cfg)
|
|
||||||
|
|
||||||
tok = AutoTokenizer.from_pretrained(MODEL_DIR)
|
tok = AutoTokenizer.from_pretrained(MODEL_DIR)
|
||||||
model = AutoModelForCausalLM.from_pretrained(MODEL_DIR, dtype=torch.float32)
|
model = AutoModelForCausalLM.from_pretrained(MODEL_DIR, dtype=torch.float32)
|
||||||
|
|||||||
@@ -25,7 +25,13 @@ cfg = SFTConfig(
|
|||||||
seed=42,
|
seed=42,
|
||||||
teacher_completions_path="data/teacher_completions_dapo1k_minimax-m3.jsonl",
|
teacher_completions_path="data/teacher_completions_dapo1k_minimax-m3.jsonl",
|
||||||
)
|
)
|
||||||
ds = load_sft_dataset(cfg)
|
ds = load_sft_dataset(
|
||||||
|
cfg.dataset_path,
|
||||||
|
cfg.dataset_split,
|
||||||
|
cfg.subset_size,
|
||||||
|
cfg.seed,
|
||||||
|
cfg.teacher_completions_path,
|
||||||
|
)
|
||||||
tok = AutoTokenizer.from_pretrained(MODEL)
|
tok = AutoTokenizer.from_pretrained(MODEL)
|
||||||
collator = SFTCollator(
|
collator = SFTCollator(
|
||||||
tok,
|
tok,
|
||||||
|
|||||||
@@ -13,7 +13,7 @@
|
|||||||
中断安全:缓存逐条落盘,重跑本脚本自动跳过已完成条目(断点续传)。
|
中断安全:缓存逐条落盘,重跑本脚本自动跳过已完成条目(断点续传)。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from ars_opd.configs import SFTConfig, TeacherGenConfig
|
from ars_opd.configs import TeacherGenConfig
|
||||||
from ars_opd.data import load_sft_dataset
|
from ars_opd.data import load_sft_dataset
|
||||||
from ars_opd.teacher import TeacherClient, generate_completions
|
from ars_opd.teacher import TeacherClient, generate_completions
|
||||||
|
|
||||||
@@ -27,15 +27,8 @@ CACHE_PATH = "data/teacher_completions_dapo1k_minimax-m3.jsonl"
|
|||||||
# 同一 seed 的洗牌序列取前缀,前 5 条与前 1000 条的头 5 条完全相同,试跑写入的
|
# 同一 seed 的洗牌序列取前缀,前 5 条与前 1000 条的头 5 条完全相同,试跑写入的
|
||||||
# 缓存在正式跑时全部命中,一分钱不浪费。
|
# 缓存在正式跑时全部命中,一分钱不浪费。
|
||||||
|
|
||||||
sft_cfg = SFTConfig(
|
|
||||||
dataset_path=DATASET_PATH,
|
|
||||||
output_dir="outputs/_unused", # 本脚本不训练,仅复用数据管线配置
|
|
||||||
subset_size=1000,
|
|
||||||
seed=42,
|
|
||||||
# teacher_completions_path 留空:此刻缓存尚不存在,取的就是 prompt-only 子集
|
# teacher_completions_path 留空:此刻缓存尚不存在,取的就是 prompt-only 子集
|
||||||
)
|
dataset = load_sft_dataset(DATASET_PATH, subset_size=1000, seed=42)
|
||||||
|
|
||||||
dataset = load_sft_dataset(sft_cfg)
|
|
||||||
prompts = [row["messages"] for row in dataset]
|
prompts = [row["messages"] for row in dataset]
|
||||||
|
|
||||||
teacher = TeacherClient(TeacherGenConfig()) # 采样参数全用 configs.py 的显式默认
|
teacher = TeacherClient(TeacherGenConfig()) # 采样参数全用 configs.py 的显式默认
|
||||||
|
|||||||
@@ -94,7 +94,13 @@ def main() -> None:
|
|||||||
|
|
||||||
# 加载顺序刻意 fail-fast:数据(毫秒级,最易配错)→ tokenizer(几 MB)→
|
# 加载顺序刻意 fail-fast:数据(毫秒级,最易配错)→ tokenizer(几 MB)→
|
||||||
# 模型(GB 级下载)。teacher 缓存缺失要在下模型之前炸出来
|
# 模型(GB 级下载)。teacher 缓存缺失要在下模型之前炸出来
|
||||||
dataset = load_sft_dataset(cfg)
|
dataset = load_sft_dataset(
|
||||||
|
cfg.dataset_path,
|
||||||
|
cfg.dataset_split,
|
||||||
|
cfg.subset_size,
|
||||||
|
cfg.seed,
|
||||||
|
cfg.teacher_completions_path,
|
||||||
|
)
|
||||||
tokenizer = AutoTokenizer.from_pretrained(STUDENT_MODEL)
|
tokenizer = AutoTokenizer.from_pretrained(STUDENT_MODEL)
|
||||||
collator = SFTCollator(
|
collator = SFTCollator(
|
||||||
tokenizer,
|
tokenizer,
|
||||||
|
|||||||
@@ -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 时的空 <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()
|
||||||
Executable
+31
@@ -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"
|
||||||
Reference in New Issue
Block a user