Files
ars-opd-rebuild/tests/test_data.py
T
iomgaa 5ea58ddf59 层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>
2026-07-18 04:51:19 -04:00

261 lines
8.7 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.
"""层 1 / T3:数据管线单测(docs/02 §5.1 规定的验证项)。
用字符级玩具 tokenizer 在 CPU 上对拍 collator 行为,不依赖网络下载真模型。
玩具模板刻意模仿 Qwen3 的关键结构:生成引导符 + no-think 时注入空思考块。
"""
import json
import pytest
from datasets import Dataset
from ars_opd.data import (
IGNORE_INDEX,
SFTCollator,
attach_teacher_completions,
prompt_key,
to_messages,
)
class ToyTokenizer:
"""字符级 tokenizer:一个字符一个 tokenid = 码点)。
模板契约与真 chat 模板同构:
- 每轮渲染成 "[role]content"
- assistant 轮(或生成引导符后)在 no-think 模式下注入 "<T></T>"(模仿
Qwen3 的空 <think>\\n\\n</think>);
- 完整渲染 == prompt 渲染 + 解答文本,保证边界可精确断言。
"""
def __init__(self, pad_token_id=0, eos_token_id=1):
self.pad_token_id = pad_token_id
self.eos_token_id = eos_token_id
def apply_chat_template(
self,
messages,
tokenize=False,
add_generation_prompt=False,
enable_thinking=False,
):
think = "" if enable_thinking else "<T></T>"
parts = []
for m in messages:
prefix = think if m["role"] == "assistant" else ""
parts.append(f"[{m['role']}]{prefix}{m['content']}")
text = "".join(parts)
if add_generation_prompt:
text += f"[assistant]{think}"
return text
def __call__(
self,
text,
truncation=False,
max_length=None,
padding=False,
add_special_tokens=False,
):
ids = [ord(c) for c in text]
if truncation and max_length is not None:
ids = ids[:max_length]
return {"input_ids": ids}
def ids_of(text):
return [ord(c) for c in text]
def row(question, answer=None):
msgs = [{"role": "user", "content": question}]
if answer is not None:
msgs.append({"role": "assistant", "content": answer})
return {"messages": msgs}
# ---------------------------------------------------------------------------
# to_messages
# ---------------------------------------------------------------------------
def test_dapo_prompt列直接归一():
ex = {"prompt": [{"role": "user", "content": "1+1=?"}], "data_source": "dapo"}
assert to_messages(ex) == {"messages": [{"role": "user", "content": "1+1=?"}]}
def test_字符串化的列表被还原():
ex = {"prompt": "[{'role': 'user', 'content': 'hi'}]"}
assert to_messages(ex)["messages"] == [{"role": "user", "content": "hi"}]
def test_坏字符串显式报错而非静默放行():
# 参考实现 except:pass 会让这行以字符串形态流进 collator
with pytest.raises(ValueError, match="无法解析"):
to_messages({"prompt": "[{'role': broken"})
def test_question列包成单user轮():
assert to_messages({"question": "2+2=?"}) == {
"messages": [{"role": "user", "content": "2+2=?"}]
}
def test_无法识别的行报错():
with pytest.raises(ValueError, match="无法识别"):
to_messages({"foo": 1})
# ---------------------------------------------------------------------------
# prompt_key(与 teacher.py 的缓存契约)
# ---------------------------------------------------------------------------
def test_同题同键_不同题不同键():
m1 = [{"role": "user", "content": "q"}]
m2 = [{"role": "user", "content": "q'"}]
assert prompt_key(m1) == prompt_key(m1)
assert prompt_key(m1) != prompt_key(m2)
def test_额外元数据字段不影响键():
plain = [{"role": "user", "content": "q"}]
noisy = [{"role": "user", "content": "q", "source": "dapo"}]
assert prompt_key(plain) == prompt_key(noisy)
# ---------------------------------------------------------------------------
# attach_teacher_completions
# ---------------------------------------------------------------------------
def write_cache(path, entries):
with open(path, "w", encoding="utf-8") as f:
for msgs, completion in entries:
f.write(
json.dumps({"key": prompt_key(msgs), "completion": completion}) + "\n"
)
def test_挂接teacher解答(tmp_path):
q = [{"role": "user", "content": "1+1=?"}]
cache = tmp_path / "cache.jsonl"
write_cache(cache, [(q, "答案是 2")])
ds = Dataset.from_list([{"messages": q}])
out = attach_teacher_completions(ds, str(cache))
assert out[0]["messages"][-1] == {"role": "assistant", "content": "答案是 2"}
def test_缓存缺键一次性报全部缺失(tmp_path):
cache = tmp_path / "cache.jsonl"
write_cache(cache, [])
ds = Dataset.from_list([row("q1"), row("q2")])
with pytest.raises(KeyError, match="2/2"):
attach_teacher_completions(ds, str(cache))
def test_自带解答的行不被覆盖(tmp_path):
cache = tmp_path / "cache.jsonl"
write_cache(cache, [])
ds = Dataset.from_list([row("q", "人写的答案")])
out = attach_teacher_completions(ds, str(cache))
assert out[0]["messages"][-1]["content"] == "人写的答案"
def test_缓存文件不存在报错():
ds = Dataset.from_list([row("q")])
with pytest.raises(FileNotFoundError):
attach_teacher_completions(ds, "/不存在/cache.jsonl")
# ---------------------------------------------------------------------------
# SFTCollator
# ---------------------------------------------------------------------------
def make_collator(**kw):
defaults = dict(max_length=1000, max_prompt_length=100, enable_thinking=False)
defaults.update(kw)
return SFTCollator(ToyTokenizer(), **defaults)
def test_基本形态_掩码与边界():
collator = make_collator()
batch = collator([row("ab", "cd")])
prompt_text = "[user]ab[assistant]<T></T>"
labels = batch["labels"][0].tolist()
# prompt 全 -100completion 位置是解答的 token
assert labels[: len(prompt_text)] == [IGNORE_INDEX] * len(prompt_text)
assert labels[len(prompt_text) :] == ids_of("cd")
assert batch["input_ids"][0].tolist() == ids_of(prompt_text + "cd")
assert batch["attention_mask"][0].tolist() == [1] * (len(prompt_text) + 2)
def test_超长解答不挤占prompt():
# 头号正确性卖点:completion 被截,prompt 一个 token 不少
prompt_text = "[user]ab[assistant]<T></T>"
collator = make_collator(max_length=len(prompt_text) + 3)
batch = collator([row("ab", "x" * 50)])
input_ids = batch["input_ids"][0].tolist()
assert input_ids[: len(prompt_text)] == ids_of(prompt_text) # prompt 完整
assert len(input_ids) == len(prompt_text) + 3 # completion 只剩预算内 3 个
def test_超长题目截断但边界不错位():
# 坑二场景:prompt 超预算被截断,completion 的 token 必须仍然精确
# (切分点用未截断长度,而非截断后长度)
collator = make_collator(max_prompt_length=10)
batch = collator([row("很长的题目" * 20, "答案")])
labels = batch["labels"][0].tolist()
non_masked = [t for t in labels if t != IGNORE_INDEX]
assert non_masked == ids_of("答案") # 解答 token 一个不错
assert sum(t == IGNORE_INDEX for t in labels) == 10 # prompt 恰被截到预算
def test_enable_thinking两种取值边界都正确():
for thinking in (False, True):
collator = make_collator(enable_thinking=thinking)
batch = collator([row("q", "ans")])
non_masked = [t for t in batch["labels"][0].tolist() if t != IGNORE_INDEX]
assert non_masked == ids_of("ans"), f"enable_thinking={thinking} 时边界错位"
def test_nothink模板注入空思考块():
# 参考实现的一次性诊断打印,在这里变成永久契约
text = ToyTokenizer().apply_chat_template(
[{"role": "user", "content": "q"}],
add_generation_prompt=True,
enable_thinking=False,
)
assert text.endswith("<T></T>")
def test_prompt_only行显式报错():
with pytest.raises(ValueError, match="prompt-only"):
make_collator()([row("没有答案的题")])
def test_左padding对齐():
collator = make_collator()
batch = collator([row("ab", "cd"), row("a", "c")])
t = batch["input_ids"].shape[1]
short_mask = batch["attention_mask"][1].tolist()
n_pad = t - short_mask.count(1)
assert n_pad > 0
assert short_mask[:n_pad] == [0] * n_pad # padding 在左
assert batch["labels"][1].tolist()[:n_pad] == [IGNORE_INDEX] * n_pad
assert batch["input_ids"][1].tolist()[:n_pad] == [0] * n_pad # pad_token_id=0
def test_pad回退到eos():
collator = SFTCollator(
ToyTokenizer(pad_token_id=None, eos_token_id=7),
max_length=100,
max_prompt_length=50,
)
assert collator.pad_token_id == 7