重构: 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:
+23
-13
@@ -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
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user