"""数据管线(IO 边缘,无论文锚点):加载 → messages 归一 → 挂接 teacher 解答 → collator。 层 1 的数据流(对应 docs/02 §1 的基线定义:SFT = 在 teacher rollout 上的离线蒸馏): DAPO parquet(prompt-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 Any import torch from datasets import Dataset, load_dataset # 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( dataset_path: str, dataset_split: str = "train", subset_size: int | None = None, seed: int = 42, teacher_completions_path: str | None = None, ) -> Dataset: """数据管线入口:加载 → 归一 → 抽子集 →(可选)挂 teacher 解答。 收散装参数而非整个 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(dataset_path, dataset_split) ds = ds.map( to_messages, remove_columns=[c for c in ds.column_names if c != "messages"], ) if subset_size is not None and subset_size < len(ds): # 非显然约束:抽子集必须在挂接 teacher 解答之前、且由 seed 完全确定—— # teacher.py 生成缓存时会走完全相同的"加载→归一→抽子集"路径,两侧 seed # 一致才能得到同一批题;否则 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 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 OPD,docs/03 §5 U3)——只渲染 prompt、 输出 prompts/prompt_attention_mask 供 model.generate 做 on-policy 生成; completion 由生成产生、labels 由 U4 的 DistillTrainer 在生成后重建,故此模式 不产 labels、也不吃 max_length。这兑现了参考实现为 on-policy 生成留的口子 (层 1 曾故意关掉,见此前 git 历史)。 与参考实现的其余差异:空 的一次性诊断打印改为单元测试断言(契约进 测试,不进运行时日志)。 """ 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 恒为 0,pad 值不参与 # 任何计算,只需要一个合法 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 SFT:messages(末轮 assistant)→ 定长张量。 返回(B = batch 大小,T = batch 内最长序列长度): - input_ids: (B, T) 左 padding - attention_mask: (B, T) padding 位置为 0 - labels: (B, T) padding 与 prompt 位置为 -100,completion 位置为 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) # 左 padding:batch 内所有序列右对齐。纯 SFT 用右 padding 也行,但左 padding # 让 trainer 能用一个标量 prompt_length 切 batch(docs/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 )