层1/T3: data.py 数据管线——加载、messages 归一、teacher 缓存挂接、双预算 collator(docs/02 §2.2-2.3)

- to_messages 三分支归一,except:pass 改显式报错(差异标注在注释)
- prompt_key: sha256 内容寻址,作为与 teacher.py 的缓存契约单点定义
- attach_teacher_completions: 缺键一次性报全,绝不静默跳过
- SFTCollator: 双预算截断 + 未截断长度定边界(坑二)+ -100 掩码 + 左 padding;
  prompt-only 行显式报错(层 2 接 on-policy 再放开)
- tests/test_data.py: 19 个单测,玩具字符级 tokenizer 覆盖超长解答/超长题目/
  enable_thinking 双取值/左 padding/缓存契约

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
2026-07-18 04:51:19 -04:00
parent 58d75cc56a
commit 5ea58ddf59
2 changed files with 588 additions and 0 deletions
+328
View File
@@ -0,0 +1,328 @@
"""数据管线(IO 边缘,无论文锚点):加载 → messages 归一 → 挂接 teacher 解答 → collator。
层 1 的数据流(对应 docs/02 §1 的基线定义:SFT = 在 teacher rollout 上的离线蒸馏):
DAPO parquetprompt-only
→ to_messages 归一成 [{"role","content"}] 列表
→ 按 seed 抽子集
→ attach_teacher_completions 从 JSONL 缓存挂上 teacher 解答(assistant 轮)
→ SFTCollator 分词、双预算截断、-100 掩码、左 padding
本模块与 teacher.py 的缓存契约由 `prompt_key` 单点定义:teacher.py 生成缓存、
本模块消费缓存,双方必须用同一个函数算键。
"""
from __future__ import annotations
import ast
import hashlib
import json
from pathlib import Path
from typing import TYPE_CHECKING, 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
# ---------------------------------------------------------------------------
# messages 归一
# ---------------------------------------------------------------------------
def _parse_stringified_list(value: str, column: str) -> list:
"""parquet 有时把 list 存成其字符串形态,用 ast 还原。
差异标注:参考实现(train_distillation.py:299-305)在这里 `except: pass` 静默吞错,
坏行会以原始字符串流进 collator,在 apply_chat_template 处以难懂的方式炸;
我们显式报错,错误信息直接指向坏数据本身。
"""
try:
parsed = ast.literal_eval(value)
except (ValueError, SyntaxError) as e:
raise ValueError(
f"{column!r} 是字符串但无法解析为 Python 字面量(坏数据行):"
f"{value[:200]!r}"
) from e
if not isinstance(parsed, (list, tuple)):
raise ValueError(f"{column!r} 解析结果不是列表:{type(parsed).__name__}")
return list(parsed)
def to_messages(example: dict[str, Any]) -> dict[str, list[dict[str, str]]]:
"""把三种来源格式归一成 messages 列:[{"role": ..., "content": ...}, ...]。
支持(与参考实现 train_distillation.py:295-321 相同的三分支):
- ``messages`` 列:直取;
- ``prompt`` 列(DAPO parquet,列名不副实——装的是完整 chat 列表):改名;
- ``question`` 列(gsm8k 风格纯文本):包成单 user 轮。
差异标注:参考实现对不认识的行 `return x` 静默放行,我们显式报错。
"""
if "messages" in example:
msgs = example["messages"]
column = "messages"
elif "prompt" in example:
msgs = example["prompt"]
column = "prompt"
elif "question" in example:
return {"messages": [{"role": "user", "content": example["question"]}]}
else:
raise ValueError(
f"无法识别的数据行:既无 messages/prompt 也无 question 列,"
f"实有列 {sorted(example.keys())}"
)
if isinstance(msgs, str):
msgs = _parse_stringified_list(msgs, column)
msgs = list(msgs)
if not msgs:
raise ValueError(f"{column!r} 是空列表(坏数据行)")
for m in msgs:
if not isinstance(m, dict) or "role" not in m or "content" not in m:
raise ValueError(
f"{column!r} 中存在非 {{role, content}} 结构的元素:{m!r}"
)
return {"messages": [{"role": m["role"], "content": m["content"]} for m in msgs]}
# ---------------------------------------------------------------------------
# teacher 解答缓存(与 teacher.py 的契约)
# ---------------------------------------------------------------------------
def prompt_key(messages: list[dict[str, str]]) -> str:
"""teacher 缓存的键:对 messages 的规范化 JSON 取 sha256。
差异标注:参考实现(distillation_trainer.py:984)用 `str(hash(prompt))`——
Python 对 str 的 hash 默认加盐,跨进程/跨次运行不稳定,缓存必然失效重生成。
sha256 内容寻址:同一道题永远同一个键。
只取 role/content 两个字段参与哈希:DAPO 行里其余元数据(data_source 等)
变了不应导致缓存失效。
"""
canon = [{"role": m["role"], "content": m["content"]} for m in messages]
return hashlib.sha256(json.dumps(canon, ensure_ascii=False).encode()).hexdigest()
def attach_teacher_completions(dataset: Dataset, jsonl_path: str) -> Dataset:
"""把 teacher 解答缓存(JSONL,每行 {"key", "completion"})挂到数据集上。
- 末轮已是 assistant 的行保持原样(数据自带解答,不覆盖);
- 任何 prompt-only 行在缓存中查不到键 → 收集齐所有缺失后一次性报错,
提示先运行 teacher 生成——绝不静默跳过(跳过 = 悄悄改变训练集组成)。
"""
path = Path(jsonl_path)
if not path.exists():
raise FileNotFoundError(
f"teacher 解答缓存不存在:{jsonl_path}。先运行 teacher.py 的批量生成。"
)
cache: dict[str, str] = {}
with open(path, encoding="utf-8") as f:
for line_no, line in enumerate(f, 1):
if not line.strip():
continue
rec = json.loads(line) # 坏行直接炸,带行号
if "key" not in rec or "completion" not in rec:
raise ValueError(f"{jsonl_path}:{line_no} 缺少 key/completion 字段")
cache[rec["key"]] = rec["completion"]
# 先整体扫描缺失,一次性报全——比在 .map 里炸第一条更省来回
missing = [
i
for i, ex in enumerate(dataset)
if ex["messages"][-1]["role"] != "assistant"
and prompt_key(ex["messages"]) not in cache
]
if missing:
raise KeyError(
f"{len(missing)}/{len(dataset)} 行在 teacher 缓存中查不到解答"
f"(首个缺失行 index={missing[0]})。检查:teacher 生成是否用了同一"
f"子集与同一 seed?(子集抽取在 load_sft_dataset 中先于挂接发生,"
f"两侧 seed 不同则键集合不同)"
)
def _attach(ex: dict[str, Any]) -> dict[str, Any]:
msgs = ex["messages"]
if msgs[-1]["role"] == "assistant":
return ex
completion = cache[prompt_key(msgs)]
return {"messages": msgs + [{"role": "assistant", "content": completion}]}
return dataset.map(_attach)
# ---------------------------------------------------------------------------
# 数据集加载(入口)
# ---------------------------------------------------------------------------
def load_sft_dataset(cfg: "SFTConfig") -> Dataset:
"""层 1 数据管线入口:加载 → 归一 → 抽子集 → 挂 teacher 解答。
返回只含 ``messages`` 一列的 Dataset,每行末轮是 assistant(可直接喂 SFTCollator)。
"""
ds = _load_raw(cfg.dataset_path, cfg.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):
# 非显然约束:抽子集必须在挂接 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)
return ds
def _load_raw(dataset_path: str, split: str) -> Dataset:
"""三分支加载:parquet 目录 / 单 parquet 文件 / HF Hub 数据集名。
差异标注:参考实现(train_distillation.py:292)对 Hub 分支硬编码 config 名
"main"(gsm8k 专用);我们不硬编码——需要特定 config 的数据集请下载成
parquet 本地加载。
"""
p = Path(dataset_path)
if p.is_dir():
return load_dataset("parquet", data_dir=dataset_path, split=split)
if dataset_path.endswith(".parquet"):
return load_dataset("parquet", data_files=dataset_path, split=split)
return load_dataset(dataset_path, split=split)
# ---------------------------------------------------------------------------
# Collator:本层最核心的一段(对拍 distillation_trainer.py:210-343
# ---------------------------------------------------------------------------
class SFTCollator:
"""把一个 batch 的 messages 变成 input_ids/attention_mask/labels。
核心设计(继承参考实现的双预算方案,docs/02 §2.3):prompt 与 completion
各自独立预算——prompt 用 max_prompt_length 截断,completion 上限是
max_length - len(截断后 prompt)。若只用一个总预算从右截断,超长解答会把
prompt 挤空,模型在"没有题目"的样本上学解答。
与参考实现的差异:
- 不支持 prompt-only 行(直接报错)。参考实现支持是为 on-policy 生成留口,
纯 SFT 下 prompt-only 行只会静默产生零 loss;层 2 接 on-policy 时再放开。
- 不返回 prompts/prompt_attention_mask(参考实现留给 vLLM 生成用,层 1 用不到)。
- 空 <think> 的一次性诊断打印改为单元测试断言(契约进测试,不进运行时日志)。
"""
def __init__(
self,
tokenizer: "Any",
max_length: int,
max_prompt_length: int,
enable_thinking: bool = False,
) -> None:
"""tokenizer 需实现 HF 接口:apply_chat_template / __call__ / pad_token_id。"""
self.tokenizer = tokenizer
self.max_length = max_length
self.max_prompt_length = max_prompt_length
self.enable_thinking = enable_thinking
# pad→eos 回退:左 padding 位置的 attention_mask 恒为 0pad 值不参与
# 任何计算,只需要一个合法 token id 占位,借用 eos 即可
if tokenizer.pad_token_id is not None:
self.pad_token_id: int = tokenizer.pad_token_id
elif tokenizer.eos_token_id is not None:
self.pad_token_id = tokenizer.eos_token_id
else:
raise ValueError("tokenizer 既无 pad_token 也无 eos_token,无法 padding")
def __call__(self, examples: list[dict[str, Any]]) -> dict[str, torch.Tensor]:
"""batch 的 messages → 定长张量。
返回(B = batch 大小,T = batch 内最长序列长度):
- input_ids: (B, T) 左 padding
- attention_mask: (B, T) padding 位置为 0
- labels: (B, T) padding 与 prompt 位置为 -100completion 位置为 token id
"""
all_input_ids: list[list[int]] = []
all_labels: list[list[int]] = []
for example in examples:
messages = example["messages"]
if len(messages) < 2 or messages[-1]["role"] != "assistant":
raise ValueError(
"SFTCollator 收到 prompt-only 行(末轮不是 assistant)。纯 SFT 下"
"它只会产生全 -100 的零 loss 样本——静默空训练。检查 teacher "
"解答是否挂接成功。"
)
# prompt = 末轮 assistant 之前的全部轮次,渲染时带生成引导符
# "<|im_start|>assistant\n..."),这样 completion 是纯解答文本的分词
formatted_prompt = self.tokenizer.apply_chat_template(
messages[:-1],
tokenize=False,
add_generation_prompt=True,
enable_thinking=self.enable_thinking,
)
# prompt 自己的预算内截断。沿用 tokenizer 默认右截断(与参考实现一致):
# 超预算的题目被截掉尾部(含生成引导符)——1024 预算下 DAPO 极少触发,
# 触发时该样本退化但不会污染边界(边界用未截断长度算,见下)
prompt_ids: list[int] = self.tokenizer(
formatted_prompt,
truncation=True,
max_length=self.max_prompt_length,
add_special_tokens=False,
)["input_ids"]
# 非显然约束(docs/02 坑一/坑二):completion 边界必须用"未截断 prompt
# 的分词长度"从整段渲染中切出。BPE 分词不满足拼接稳定性,分开渲染
# prompt 和 completion 再拼接 ≠ 整段渲染后分词;而若用截断后长度当切分
# 点,会把 prompt 尾部的 token 误标成 completion——静默的语义错误。
formatted_full = self.tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=False,
enable_thinking=self.enable_thinking,
)
full_ids: list[int] = self.tokenizer(
formatted_full, truncation=False, add_special_tokens=False
)["input_ids"]
untruncated_prompt_len = len(
self.tokenizer(
formatted_prompt, truncation=False, add_special_tokens=False
)["input_ids"]
)
completion_ids = full_ids[untruncated_prompt_len:]
# completion 预算 = 总预算 - 截断后 prompt 实长。配置校验
# max_prompt_length < max_length)保证它恒 > 0
completion_budget = self.max_length - len(prompt_ids)
completion_ids = completion_ids[:completion_budget]
all_input_ids.append(prompt_ids + completion_ids)
# prompt 位置标 -100:题目不产生 loss,只学解答
all_labels.append([IGNORE_INDEX] * len(prompt_ids) + completion_ids)
# 左 paddingbatch 内所有序列右对齐。纯 SFT 用右 padding 也行,但左 padding
# 让 trainer 能用一个标量 prompt_length 切 batchdocs/02 §2.4),且与
# 层 2+ 的生成场景(生成必须左 padding)统一,全项目只有一种 padding 约定
return {
"input_ids": _left_pad(all_input_ids, self.pad_token_id), # (B, T)
"attention_mask": _left_pad(
[[1] * len(ids) for ids in all_input_ids], 0
), # (B, T)
"labels": _left_pad(all_labels, IGNORE_INDEX), # (B, T)
}
def _left_pad(seqs: list[list[int]], pad_value: int) -> torch.Tensor:
"""把变长序列在左侧补齐成 (B, T) 张量,T = batch 内最大长度。"""
t_max = max(len(s) for s in seqs)
return torch.tensor(
[[pad_value] * (t_max - len(s)) + s for s in seqs], dtype=torch.long
)