层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>
This commit is contained in:
2026-07-18 07:52:53 -04:00
parent f5bb852fde
commit 20e6d97427
7 changed files with 414 additions and 7 deletions
+143
View File
@@ -0,0 +1,143 @@
"""层 1 / T2:teacher 批量生成与缓存单测。
用假 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")