重构: load_sft_dataset 改吃散装参数(磨平接口回看记录的毛刺)

深模块修正:本函数只用 5 个字段,却索要整个 SFTConfig——层 1 无痛,但诊断脚本
被迫伪造 output_dir(4 处 /tmp/diag、outputs/_unused),层 2 更因 DistillConfig
无 teacher_completions_path 而无法复用。改收 dataset_path/split/subset_size/seed/
teacher_completions_path 五个散装参数(接口终于比实现轻)。

- data.py: 签名改散装参数;移除 TYPE_CHECKING 的 SFTConfig 依赖
- train_sft / diag_loss_probe / diag_collator: 仍持 SFTConfig(喂 collator),改调用点
- diag_generate / generate_teacher_completions: 只为 load 而造 config,直接丢弃、
  去掉伪造 output_dir,改传字面量
- 为 U5 层 2 训练脚本能直接 load_sft_dataset(distill_cfg 的字段) 铺路

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-07-19 05:21:16 -04:00
parent 0ca60ea93f
commit 404abc22bf
6 changed files with 50 additions and 34 deletions
+23 -13
View File
@@ -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)。这样层 1SFTConfig)、层 2DistillConfig,无 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
+7 -1
View File
@@ -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"]
+3 -8
View File
@@ -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)
+7 -1
View File
@@ -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,
+3 -10
View File
@@ -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( # teacher_completions_path 留空:此刻缓存尚不存在,取的就是 prompt-only 子集
dataset_path=DATASET_PATH, dataset = load_sft_dataset(DATASET_PATH, subset_size=1000, seed=42)
output_dir="outputs/_unused", # 本脚本不训练,仅复用数据管线配置
subset_size=1000,
seed=42,
# teacher_completions_path 留空:此刻缓存尚不存在,取的就是 prompt-only 子集
)
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 的显式默认
+7 -1
View File
@@ -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,