From 404abc22bf4cdc22392fc07e015b801e895fc39c Mon Sep 17 00:00:00 2001 From: iomgaa Date: Sun, 19 Jul 2026 05:21:16 -0400 Subject: [PATCH] =?UTF-8?q?=E9=87=8D=E6=9E=84:=20load=5Fsft=5Fdataset=20?= =?UTF-8?q?=E6=94=B9=E5=90=83=E6=95=A3=E8=A3=85=E5=8F=82=E6=95=B0=EF=BC=88?= =?UTF-8?q?=E7=A3=A8=E5=B9=B3=E6=8E=A5=E5=8F=A3=E5=9B=9E=E7=9C=8B=E8=AE=B0?= =?UTF-8?q?=E5=BD=95=E7=9A=84=E6=AF=9B=E5=88=BA=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 深模块修正:本函数只用 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) --- ars_opd/data.py | 36 ++++++++++++++++--------- scripts/diag_collator.py | 8 +++++- scripts/diag_generate.py | 11 +++----- scripts/diag_loss_probe.py | 8 +++++- scripts/generate_teacher_completions.py | 13 +++------ scripts/train_sft.py | 8 +++++- 6 files changed, 50 insertions(+), 34 deletions(-) diff --git a/ars_opd/data.py b/ars_opd/data.py index 2e861ee..24b8962 100644 --- a/ars_opd/data.py +++ b/ars_opd/data.py @@ -18,14 +18,11 @@ import ast import hashlib import json from pathlib import Path -from typing import TYPE_CHECKING, Any +from typing import Any import torch from datasets import Dataset, load_dataset -if TYPE_CHECKING: - from ars_opd.configs import SFTConfig - # F.cross_entropy 的 ignore_index 默认值;标了它的位置不产生 loss IGNORE_INDEX = -100 @@ -165,23 +162,36 @@ def attach_teacher_completions(dataset: Dataset, jsonl_path: str) -> Dataset: # --------------------------------------------------------------------------- -def load_sft_dataset(cfg: "SFTConfig") -> Dataset: - """层 1 数据管线入口:加载 → 归一 → 抽子集 → 挂 teacher 解答。 +def load_sft_dataset( + 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( to_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.py 生成缓存时会走完全相同的"加载→归一→抽子集"路径,两侧 seed - # 一致才能得到同一批题;否则 attach 处大面积缓存 miss 报错。 - ds = ds.shuffle(seed=cfg.seed).select(range(cfg.subset_size)) - if cfg.teacher_completions_path is not None: - ds = attach_teacher_completions(ds, cfg.teacher_completions_path) + # 一致才能得到同一批题;否则 attach 处大面积缓存 miss 报错。层 2 与层 1 + # 用同 seed 同 subset_size,才能在同一批题上对比 SFT 与蒸馏。 + ds = ds.shuffle(seed=seed).select(range(subset_size)) + if teacher_completions_path is not None: + ds = attach_teacher_completions(ds, teacher_completions_path) return ds diff --git a/scripts/diag_collator.py b/scripts/diag_collator.py index e5565f2..d3607aa 100644 --- a/scripts/diag_collator.py +++ b/scripts/diag_collator.py @@ -46,7 +46,13 @@ def main() -> None: seed=42, 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"] completion_text = msgs[-1]["content"] diff --git a/scripts/diag_generate.py b/scripts/diag_generate.py index cea06e9..4605b47 100644 --- a/scripts/diag_generate.py +++ b/scripts/diag_generate.py @@ -7,19 +7,14 @@ import torch from transformers import AutoModelForCausalLM, AutoTokenizer -from ars_opd.configs import SFTConfig from ars_opd.data import load_sft_dataset MODEL_DIR = "/data/zym/outputs/sft_qwen3-0.6b_dapo1k" # 正式 1 epoch 的产物 -cfg = SFTConfig( - dataset_path="data/dapo-math-17k-unique.parquet", - output_dir="/tmp/diag", - subset_size=1000, - seed=42, - # 不挂 teacher 解答:只取题目做推理输入 +# 不挂 teacher 解答(teacher_completions_path 缺省):只取题目做推理输入 +ds = load_sft_dataset( + "data/dapo-math-17k-unique.parquet", subset_size=1000, seed=42 ) -ds = load_sft_dataset(cfg) tok = AutoTokenizer.from_pretrained(MODEL_DIR) model = AutoModelForCausalLM.from_pretrained(MODEL_DIR, dtype=torch.float32) diff --git a/scripts/diag_loss_probe.py b/scripts/diag_loss_probe.py index 1094c25..45f75f5 100644 --- a/scripts/diag_loss_probe.py +++ b/scripts/diag_loss_probe.py @@ -25,7 +25,13 @@ cfg = SFTConfig( seed=42, 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) collator = SFTCollator( tok, diff --git a/scripts/generate_teacher_completions.py b/scripts/generate_teacher_completions.py index a7fe366..6f46f0d 100644 --- a/scripts/generate_teacher_completions.py +++ b/scripts/generate_teacher_completions.py @@ -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.teacher import TeacherClient, generate_completions @@ -27,15 +27,8 @@ CACHE_PATH = "data/teacher_completions_dapo1k_minimax-m3.jsonl" # 同一 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 子集 -) - -dataset = load_sft_dataset(sft_cfg) +# teacher_completions_path 留空:此刻缓存尚不存在,取的就是 prompt-only 子集 +dataset = load_sft_dataset(DATASET_PATH, subset_size=1000, seed=42) prompts = [row["messages"] for row in dataset] teacher = TeacherClient(TeacherGenConfig()) # 采样参数全用 configs.py 的显式默认 diff --git a/scripts/train_sft.py b/scripts/train_sft.py index 874ac95..f2fa01b 100644 --- a/scripts/train_sft.py +++ b/scripts/train_sft.py @@ -94,7 +94,13 @@ def main() -> None: # 加载顺序刻意 fail-fast:数据(毫秒级,最易配错)→ tokenizer(几 MB)→ # 模型(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) collator = SFTCollator( tokenizer,