20e6d97427
- 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>
144 lines
5.2 KiB
Python
144 lines
5.2 KiB
Python
"""层 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")
|