层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:
@@ -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")
|
||||
Reference in New Issue
Block a user