Files
ars-opd-rebuild/ars_opd/data.py
T
iomgaa e5a28e8e77 层2/U3: SFTCollator 放开 prompt-only 模式(供 on-policy 生成)
data.py(对应 docs/03 §5 U3):
- 加 prompt_only 开关:True 时输出 prompts/prompt_attention_mask(不产 labels,
  由 U4 生成后重建);False 时 SFT 双预算路径逐字不变
- max_length 改可选:prompt-only 无总预算;SFT 模式缺它构造即报错
- 兑现 T3 为 on-policy 生成预留的口子;生成用左 padding(右边界对齐)

test_data.py:
- 新增 prompt_only 模式:返回 prompt 张量/不报错、左 padding、截断、剥末轮 assistant
- 回归守卫:SFT 模式仍拒绝 prompt-only 行("SFT 路径行为不变")

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-19 05:02:24 -04:00

398 lines
18 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""数据管线(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 变成训练/生成所需的定长张量,两种模式二选一。
prompt_only=False(层 1 SFT,默认)——输出 input_ids/attention_mask/labels
核心设计(继承参考实现的双预算方案,docs/02 §2.3):prompt 与 completion
各自独立预算——prompt 用 max_prompt_length 截断,completion 上限是
max_length - len(截断后 prompt)。若只用一个总预算从右截断,超长解答会把
prompt 挤空,模型在"没有题目"的样本上学解答。要求每行末轮是 assistant,
否则报错(prompt-only 行在纯 SFT 下只产生零 loss = 静默空训练)。
prompt_only=True(层 2 white-box OPDdocs/03 §5 U3)——只渲染 prompt、
输出 prompts/prompt_attention_mask 供 model.generate 做 on-policy 生成;
completion 由生成产生、labels 由 U4 的 DistillTrainer 在生成后重建,故此模式
不产 labels、也不吃 max_length。这兑现了参考实现为 on-policy 生成留的口子
(层 1 曾故意关掉,见此前 git 历史)。
与参考实现的其余差异:空 <think> 的一次性诊断打印改为单元测试断言(契约进
测试,不进运行时日志)。
"""
def __init__(
self,
tokenizer: "Any",
max_prompt_length: int,
max_length: int | None = None,
enable_thinking: bool = False,
prompt_only: bool = False,
) -> None:
"""tokenizer 需实现 HF 接口:apply_chat_template / __call__ / pad_token_id。
max_length 仅 SFT 模式需要(completion 预算依赖它);prompt_only 模式下
completion 是生成的、无总预算,故 max_length 可为 None。
"""
if not prompt_only and max_length is None:
raise ValueError(
"SFT 模式(prompt_only=False)必须提供 max_length——completion "
"预算 = max_length - len(prompt),缺它无法确定解答截断点。"
)
self.tokenizer = tokenizer
self.max_length = max_length
self.max_prompt_length = max_prompt_length
self.enable_thinking = enable_thinking
self.prompt_only = prompt_only
# 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]:
"""按模式分派:prompt_only 走生成用 prompt 张量,否则走 SFT 双预算。"""
if self.prompt_only:
return self._collate_prompt_only(examples)
return self._collate_sft(examples)
def _collate_prompt_only(
self, examples: list[dict[str, Any]]
) -> dict[str, torch.Tensor]:
"""层 2:只渲染 prompt 供 on-policy 生成,不产 completion/labels。
返回(B = batch 大小,P = batch 内最长 prompt 长度):
- prompts: (B, P) 左 padding
- prompt_attention_mask: (B, P) padding 位置为 0
非显然约束:生成必须左 padding——所有 prompt 右对齐到同一右边界,
model.generate 从该边界统一续写;右 padding 会让短 prompt 的生成从 pad
中间开始,全乱。这也是层 1 SFT 就选左 padding 的原因(全项目一种约定)。
"""
all_prompt_ids: list[list[int]] = []
for example in examples:
messages = example["messages"]
# prompt-only 数据末轮是 user;若末轮已是 assistant 则剥掉,取生成前上下文
prompt_msgs = (
messages[:-1] if messages[-1]["role"] == "assistant" else messages
)
if not prompt_msgs:
raise ValueError(
"prompt_only collator 收到空 prompt(无可生成的上下文)"
)
# 与 SFT 模式同样带生成引导符渲染(add_generation_prompt=True):
# prompt 末尾就是 "<|im_start|>assistant\n...",生成从此续写
formatted_prompt = self.tokenizer.apply_chat_template(
prompt_msgs,
tokenize=False,
add_generation_prompt=True,
enable_thinking=self.enable_thinking,
)
prompt_ids: list[int] = self.tokenizer(
formatted_prompt,
truncation=True,
max_length=self.max_prompt_length,
add_special_tokens=False,
)["input_ids"]
all_prompt_ids.append(prompt_ids)
return {
"prompts": _left_pad(all_prompt_ids, self.pad_token_id), # (B, P)
"prompt_attention_mask": _left_pad(
[[1] * len(ids) for ids in all_prompt_ids], 0
), # (B, P)
}
def _collate_sft(self, examples: list[dict[str, Any]]) -> dict[str, torch.Tensor]:
"""层 1 SFTmessages(末轮 assistant)→ 定长张量。
返回(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
)