层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>
This commit is contained in:
+79
-10
@@ -206,32 +206,48 @@ def _load_raw(dataset_path: str, split: str) -> Dataset:
|
||||
|
||||
|
||||
class SFTCollator:
|
||||
"""把一个 batch 的 messages 变成 input_ids/attention_mask/labels。
|
||||
"""把一个 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 挤空,模型在"没有题目"的样本上学解答。
|
||||
prompt 挤空,模型在"没有题目"的样本上学解答。要求每行末轮是 assistant,
|
||||
否则报错(prompt-only 行在纯 SFT 下只产生零 loss = 静默空训练)。
|
||||
|
||||
与参考实现的差异:
|
||||
- 不支持 prompt-only 行(直接报错)。参考实现支持是为 on-policy 生成留口,
|
||||
纯 SFT 下 prompt-only 行只会静默产生零 loss;层 2 接 on-policy 时再放开。
|
||||
- 不返回 prompts/prompt_attention_mask(参考实现留给 vLLM 生成用,层 1 用不到)。
|
||||
- 空 <think> 的一次性诊断打印改为单元测试断言(契约进测试,不进运行时日志)。
|
||||
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 历史)。
|
||||
|
||||
与参考实现的其余差异:空 <think> 的一次性诊断打印改为单元测试断言(契约进
|
||||
测试,不进运行时日志)。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: "Any",
|
||||
max_length: int,
|
||||
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。"""
|
||||
"""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:
|
||||
@@ -242,7 +258,60 @@ class SFTCollator:
|
||||
raise ValueError("tokenizer 既无 pad_token 也无 eos_token,无法 padding")
|
||||
|
||||
def __call__(self, examples: list[dict[str, Any]]) -> dict[str, torch.Tensor]:
|
||||
"""batch 的 messages → 定长张量。
|
||||
"""按模式分派: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
|
||||
|
||||
Reference in New Issue
Block a user