Files
iomgaa 20e6d97427 层1/T2: teacher.py 批量生成 + sha256 JSONL 缓存;teacher 改定 MiniMax-M3
- teacher.py: 通用 OpenAI 兼容客户端(配置驱动 base_url,替代 OpenRouter 专用);
  缓存即断点(逐条落盘+flush,重跑自动续传);单条失败先落盘其余、结束汇总显式报错;
  M3 思考段 <think>...</think> 入库前剥离(只剥开头一段)
- configs.py: 新增 TeacherGenConfig(采样参数显式化;连接三元组走 .env)
- scripts/generate_teacher_completions.py: 自包含生成脚本(本地跑,与训练侧
  同 seed 同子集约束已注明)
- teacher 决策变更同步:.env.example / docs/00 关键设定与存档点 / docs/02
- tests/test_teacher.py: 10 个单测(假客户端注入),含与 attach 的端到端契约闭环

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-18 07:52:53 -04:00

144 lines
5.2 KiB
Python
Raw Permalink 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 / T2teacher 批量生成与缓存单测。
用假 OpenAI 客户端注入(TeacherClient 的测试口),验证思考段剥离、缓存契约
(与 data.attach_teacher_completions 的端到端闭环)、断点续传、失败汇总。
"""
from types import SimpleNamespace
import pytest
from datasets import Dataset
from ars_opd.configs import TeacherGenConfig
from ars_opd.data import attach_teacher_completions, prompt_key
from ars_opd.teacher import TeacherClient, _load_teacher_env, generate_completions
class FakeClient:
"""最小 OpenAI 客户端替身:chat.completions.create 按 responder 出内容。"""
def __init__(self, responder):
self.calls = []
self._responder = responder
self.chat = SimpleNamespace(completions=SimpleNamespace(create=self._create))
def _create(self, model, messages, **kwargs):
self.calls.append(messages)
content = self._responder(messages)
return SimpleNamespace(
choices=[SimpleNamespace(message=SimpleNamespace(content=content))]
)
def make_teacher(responder, **cfg_overrides):
cfg = TeacherGenConfig(**cfg_overrides)
return TeacherClient(cfg, client=FakeClient(responder), model="fake-m3")
def user(q):
return [{"role": "user", "content": q}]
# ---------------------------------------------------------------------------
# TeacherClient.generate
# ---------------------------------------------------------------------------
def test_剥离开头思考段():
teacher = make_teacher(lambda m: "<think>心算一下</think>\n答案是 42")
assert teacher.generate(user("q")) == "答案是 42"
def test_正文中的think字样不误删():
teacher = make_teacher(lambda m: "<think>x</think>正文提到 <think> 标签本身")
assert teacher.generate(user("q")) == "正文提到 <think> 标签本身"
def test_只剩思考段等于空解答_报错():
teacher = make_teacher(lambda m: "<think>思考被截断在半途")
# 未闭合的 think 段剥不掉,但闭合后为空的要报错
teacher_empty = make_teacher(lambda m: "<think>只有思考</think> ")
with pytest.raises(ValueError, match="空解答"):
teacher_empty.generate(user("q"))
# 未闭合时保留原文(宁可保留可疑内容也不静默删成空)
assert "<think>" in teacher.generate(user("q"))
def test_关闭strip_think则原样保留():
teacher = make_teacher(lambda m: "<think>a</think>b", strip_think=False)
assert teacher.generate(user("q")) == "<think>a</think>b"
def test_system_prompt前置():
teacher = make_teacher(lambda m: "ok", system_prompt="你是数学助教")
teacher.generate(user("q"))
sent = teacher.client.calls[0]
assert sent[0] == {"role": "system", "content": "你是数学助教"}
assert sent[1]["role"] == "user"
def test_注入client但不给model报错():
with pytest.raises(ValueError, match="model"):
TeacherClient(TeacherGenConfig(), client=FakeClient(lambda m: "x"), model=None)
# ---------------------------------------------------------------------------
# generate_completions:缓存契约与断点续传
# ---------------------------------------------------------------------------
def test_端到端契约_生成的缓存能被attach消费(tmp_path):
cache = str(tmp_path / "cache.jsonl")
prompts = [user("1+1=?"), user("2+2=?")]
teacher = make_teacher(lambda m: f"对「{m[-1]['content']}」的解答")
generate_completions(prompts, cache, teacher)
ds = Dataset.from_list([{"messages": p} for p in prompts])
out = attach_teacher_completions(ds, cache)
assert out[0]["messages"][-1]["content"] == "对「1+1=?」的解答"
assert out[1]["messages"][-1]["content"] == "对「2+2=?」的解答"
def test_断点续传_已缓存的不重新生成(tmp_path):
cache = str(tmp_path / "cache.jsonl")
prompts = [user("q1"), user("q2")]
teacher = make_teacher(lambda m: "a")
generate_completions([prompts[0]], cache, teacher)
assert len(teacher.client.calls) == 1
generate_completions(prompts, cache, teacher) # q1 命中缓存
assert len(teacher.client.calls) == 2 # 只多了 q2 一次调用
def test_单条失败_其余落盘_结束时汇总报错(tmp_path):
cache = str(tmp_path / "cache.jsonl")
prompts = [user("好题"), user("坏题")]
def responder(m):
if m[-1]["content"] == "坏题":
raise RuntimeError("网关 500")
return "解答"
teacher = make_teacher(responder)
with pytest.raises(RuntimeError, match="1/2"):
generate_completions(prompts, cache, teacher)
# 成功的那条已经在缓存里,重跑只会补坏题
from ars_opd.teacher import _cached_keys
from pathlib import Path
assert _cached_keys(Path(cache)) == {prompt_key(prompts[0])}
# ---------------------------------------------------------------------------
# .env 读取
# ---------------------------------------------------------------------------
def test_env缺失显式报错(monkeypatch):
for name in ("TEACHER_API_BASE", "TEACHER_API_KEY", "TEACHER_MODEL"):
monkeypatch.delenv(name, raising=False)
with pytest.raises(ValueError, match="TEACHER_API_BASE"):
_load_teacher_env(env_file="/不存在的路径/.env")