Files
Video-Tree-TRM5/tests/integration/test_v3_contract_e2e.py
T
iomgaa 8958eee11b refactor: remove block-sequential gate path and gate_block knob (algo #6)
config/train_videomme.yaml 同时收录待入库的实验配置变更(run_id v2 /
concurrency 32 / batch_size 40)。tests/integration/test_v3_contract_e2e.py
的 run_id 断言按 Task 5 显式契约同步修正(原断言依赖旧隐式实例注入)。
2026-07-17 04:40:30 -04:00

380 lines
16 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.
"""question-gen v3 Phase 1 契约全链路集成回归(Task 13)。
把 T3-T12 建立的 QuestionUnit(single|pair) 契约在 build_pools → build_batches →
run_inference 全链路上做一次端到端回归护栏。构造「混格题库」(若干 AR pair + 若干
非 AR single)真实驱动三个生产入口,逐条锁定六条契约:
1. pair 不拆:混格题库经 build_pools 后,任一 pair 的 original+mirror 落在同一池。
2. 同批:诊断池经 build_batches 后,任一 pair 的两题在同一个 batch。
3. 双向 AND 聚合:pair P 对 Q 错 → 该 unit 计错;P 对 Q 对 → 计对。
4. unit 粒度 total/correctN single + M pair 的库,total == N + Mpair 计 1)。
5. 孤儿被剔:只有 original 没有 mirror 的悬挂 pair → 聚合剔除 + 告警、不计入 total。
6. 非 AR byte-identical:纯 single 题库 build_batches 输出与旧逐题黄金参照逐字节一致。
黄金参照 helper_reference_build_batches / _ids)与最小构造 helper_single /
_pair)复用 tests/unit/test_non_ar_byte_identical.py,不重造。fake 只在 LLM/工具/
prompt 外部依赖层,build_pools/build_batches/run_inference/聚合逻辑全部真跑。
"""
from __future__ import annotations
import json
from typing import Any
from unittest.mock import AsyncMock
import pytest
from loguru import logger
from app.harness.batching import build_batches
from app.harness.inference import InferenceResult, run_inference
from app.harness.log import HarnessLog
from app.harness.pools import build_pools
from core.types import GeneratedQuestion, LLMResponse
from tests.unit.test_non_ar_byte_identical import (
_ids,
_pair,
_reference_build_batches,
_single,
)
# ── 外部依赖 fake(仅 LLM/工具/prompt,被测链路真跑) ──────────────────
def _make_question(
question_id: str,
*,
task_type: str = "RETRIEVAL",
answer: str = "B",
pair_id: str | None = None,
question_role: str = "single",
flip_axis: str | None = None,
video_id: str = "v1",
) -> GeneratedQuestion:
"""构造混格题目;pair_id 非空时视为孪生对成员(共享 flip_axis)。"""
return GeneratedQuestion(
question_id=question_id,
video_id=video_id,
task_type=task_type,
question="测试问题",
options=("A. 选项A", "B. 选项B", "C. 选项C", "D. 选项D"),
answer=answer,
source_nodes=("L1_001",),
difficulty="medium",
pair_id=pair_id,
question_role=question_role,
flip_axis=flip_axis if pair_id else None,
)
def _make_llm_response(answer: str = "B") -> LLMResponse:
"""构造 submit_answer 场景的 LLMResponse(固定预测某答案)。"""
content = json.dumps(
{
"reflect": {"observation": "找到答案"},
"plan": {"next_step": "提交"},
"action": {
"tool": "submit_answer",
"args": {"answer": answer, "evidence": "证据", "reasoning": "推理"},
},
}
)
return LLMResponse(
content=content,
thinking="思考",
model="test-model",
provider="test",
prompt_tokens=100,
completion_tokens=50,
latency_ms=200,
ttft_ms=30.0,
max_inter_token_ms=5.0,
cache_hit=False,
call_id="test-call-001",
)
async def _stub_tool_dispatch(
tool_name: str, args: dict[str, Any], *, context: dict[str, Any]
) -> str:
"""测试用工具调度:只认 submit_answer,未知工具 fail-loud。"""
if tool_name == "submit_answer":
return "答案已提交"
raise ValueError(f"未知工具: {tool_name}")
def _stub_prompt_builder(qa: GeneratedQuestion) -> tuple[str, str]:
"""测试用 prompt 构建。"""
return "系统提示词", f"用户问题: {qa.question}"
@pytest.fixture
def harness_log(tmp_path: Any, request: Any) -> HarnessLog:
"""临时 HarnessLog 实例(run_id 标记为 test-run)。"""
db_path = str(tmp_path / f"harness_{id(request)}.db")
log = HarnessLog(db_path, "test-run")
yield log
log.close()
# ── 断言 helper(把跨池/跨批的成员定位与 pair 不拆判定收口一处) ────────
def _membership(labeled_groups: list[tuple[str, list[GeneratedQuestion]]]) -> dict[str, str]:
"""把带标签的分组(池或 batch)展平成 question_id → 组标签,并断言组间互斥。
参数:
labeled_groups: (标签, 题目列表) 列表;标签即池名或 batch 下标字符串。
返回:
question_id → 所属组标签映射。
关键实现:
同一 question_id 落入两个组即 fail(三池互斥 / 单题不重复入批的前提被破坏)。
"""
location: dict[str, str] = {}
for label, group in labeled_groups:
for q in group:
assert q.question_id not in location, (
f"{q.question_id} 同属 {location.get(q.question_id)}{label}(组未互斥)"
)
location[q.question_id] = label
return location
def _pair_together(location: dict[str, str], pair_id: str, *, allow_absent: bool) -> bool:
"""判定某 pair 的 original+mirror 是否整锁同组,返回是否被采样进组。
参数:
location: question_id → 组标签(来自 _membership)。
pair_id: 孪生对标识(成员 id 约定为 f"{pair_id}_o" / f"{pair_id}_m")。
allow_absent: True 时允许整对都未进组(build_pools 场景,采样可不命中);
False 时要求必进组(build_batches 全错场景,必进批)。
返回:
该对是否被完整采样进某组(用于调用方统计非空护栏)。
异常:
AssertionError: 半只进组(被拆)或两题分属不同组(跨组)。
"""
o_loc, m_loc = location.get(f"{pair_id}_o"), location.get(f"{pair_id}_m")
if allow_absent and o_loc is None and m_loc is None:
return False
assert o_loc is not None and m_loc is not None, (
f"pair {pair_id} 被拆:original@{o_loc} mirror@{m_loc}(半只进组)"
)
assert o_loc == m_loc, f"pair {pair_id} 跨组:original@{o_loc} mirror@{m_loc}"
return True
async def _run_capturing_warnings(
questions: list[GeneratedQuestion], log: HarnessLog
) -> tuple[InferenceResult, list[str]]:
"""用固定预测 "B" 的 fake LLM 真跑 run_inference,并捕获 WARNING 日志。
参数:
questions: 待推理的混格题库。
log: HarnessLog 实例。
返回:
(InferenceResult, 捕获到的 WARNING 消息列表)。
"""
llm = AsyncMock()
llm.chat.return_value = _make_llm_response(answer="B")
captured: list[str] = []
sink_id = logger.add(captured.append, level="WARNING", format="{message}")
try:
result = await run_inference(
questions,
llm=llm,
tool_dispatch_fn=_stub_tool_dispatch,
prompt_builder=_stub_prompt_builder,
log=log,
run_id="run-v3-contract",
concurrency=4,
max_steps=10,
skill_mode="auto",
)
finally:
logger.remove(sink_id)
return result, captured
# ── 契约 1:pair 全程不拆,两题永不跨池 ───────────────────────────────
class TestPairNeverSplitAcrossPools:
"""混格题库经 build_pools 后,任一 pair 的 original+mirror 落在同一池。"""
def _mixed_benchmark(self) -> tuple[list[GeneratedQuestion], dict[str, bool]]:
"""24 non-AR single + 8 AR pair 混格库,附单元级基线对错。
single:偶数下标对、奇数下标错(12 对 12 错)。
pair:前 4 对两题皆对(单元对),后 4 对 original 错(单元错)。
"""
singles = [_single(f"s{i}", task_type="RETRIEVAL") for i in range(24)]
correctness: dict[str, bool] = {q.question_id: (i % 2 == 0) for i, q in enumerate(singles)}
pair_questions: list[GeneratedQuestion] = []
for i in range(8):
original, mirror = _pair(f"p{i}", task_type="AR")
pair_questions.extend([original, mirror])
correctness[original.question_id] = i < 4
correctness[mirror.question_id] = True
return singles + pair_questions, correctness
def test_no_pair_lands_in_two_pools(self) -> None:
"""任一 pair 两题要么同池、要么都未被采样,绝不分属不同池。"""
questions, correctness = self._mixed_benchmark()
pools = build_pools(
questions,
correctness,
diag_cfg={
"size": 10,
"correct_ratio": 0.5,
"task_types": None,
"seed": 1,
"min_per_class": None,
},
val_cfg={
"size": 8,
"correct_ratio": 0.5,
"task_types": None,
"seed": 1,
"min_per_class": None,
},
test_cfg={"size": 6, "seed": 1},
baseline_run_id="baseline",
)
pool_of = _membership(
[
("diagnosis", pools.diagnosis),
("validation", pools.validation),
("test", pools.test),
]
)
sampled = sum(_pair_together(pool_of, f"p{i}", allow_absent=True) for i in range(8))
# 非空护栏:确保确有 pair 被采样进池,断言不是空转
assert sampled > 0, "无 pair 进入任一池,pair-不拆断言未被实质覆盖"
# ── 契约 2:pair 整锁同批 ─────────────────────────────────────────────
class TestPairStaysInSameBatch:
"""诊断池经 build_batches 后,任一 pair 的两题落在同一个 batch。"""
def _all_wrong_diagnosis(self) -> tuple[list[GeneratedQuestion], dict[str, bool]]:
"""全错混格诊断池:6 non-AR single + 4 AR pair,皆错以确保全部进批。"""
singles = [_single(f"s{i}", task_type="RETRIEVAL") for i in range(6)]
correctness: dict[str, bool] = {q.question_id: False for q in singles}
items: list[GeneratedQuestion] = list(singles)
for i in range(4):
original, mirror = _pair(f"p{i}", task_type="AR")
items.extend([original, mirror])
correctness[original.question_id] = False
correctness[mirror.question_id] = False
return items, correctness
def test_pair_members_share_batch(self) -> None:
"""全错混格诊断池分批后,每对两题同批(pair 占 2 容量整锁不拆)。"""
items, correctness = self._all_wrong_diagnosis()
batches, selected = build_batches(
items, correctness, batch_size=6, min_class_per_batch=2, seed=3, correct_ratio=0.0
)
assert selected == len(items), "全错单元应全部进批(selected == 展开题数)"
batch_of = _membership([(str(idx), batch) for idx, batch in enumerate(batches)])
in_batches = sum(_pair_together(batch_of, f"p{i}", allow_absent=False) for i in range(4))
assert in_batches == 4, "并非全部 pair 都进批,同批断言未实质覆盖"
# ── 契约 3+4+5run_inference 双向 AND + unit 粒度 total + 孤儿剔除 ────
class TestInferenceUnitAggregationEndToEnd:
"""混格题库真跑 run_inference:双向 AND、unit total/correct、孤儿剔除+告警。"""
def _mixed_questions(self) -> list[GeneratedQuestion]:
"""3 single(皆对)+ pairA(两题皆对)+ pairBP 对 Q 错)+ 孤儿 pair。
LLM 固定预测 "B"answer=="B" 即对、answer=="A" 即错,从而在单次推理里
制造出 unit 级双向 AND 的对/错两种结果。
"""
def _pair_member(qid: str, pid: str, role: str, answer: str) -> GeneratedQuestion:
return _make_question(
qid,
task_type="AR",
answer=answer,
pair_id=pid,
question_role=role,
flip_axis="before_after",
)
return [
_make_question("s1", answer="B"),
_make_question("s2", answer="B"),
_make_question("s3", answer="B"),
# pairA:两题皆答对 → 双向 AND 判对
_pair_member("pa_o", "pa", "pair_original", "B"),
_pair_member("pa_m", "pa", "pair_mirror", "B"),
# pairBoriginal 对、mirror 错(answer=A)→ 双向 AND 判错
_pair_member("pb_o", "pb", "pair_original", "B"),
_pair_member("pb_m", "pb", "pair_mirror", "A"),
# 孤儿 pair:只有 original,没有 mirror → 聚合剔除 + 告警
_pair_member("orphan_o", "orphan_p", "pair_original", "B"),
]
def _assert_task_type_breakdown(self, result: InferenceResult) -> None:
"""契约 3 细化到 task_typeAR 两对仅 pairA 对,RETRIEVAL 三题皆对。"""
assert result.per_task_type["AR"]["total"] == 2
assert result.per_task_type["AR"]["correct"] == 1
assert result.per_task_type["RETRIEVAL"]["total"] == 3
assert result.per_task_type["RETRIEVAL"]["correct"] == 3
def _assert_all_persisted(self, log: HarnessLog, questions: list[GeneratedQuestion]) -> None:
"""逐题溯源保留:含被剔除的孤儿题在内,每题仍逐题落 predictions。"""
rows = log.query("SELECT * FROM predictions WHERE run_id = ?", ("run-v3-contract",))
persisted = {r["question_id"] for r in rows}
assert "orphan_o" in persisted, "孤儿题未逐题落库(逐题溯源被破坏)"
assert persisted == {q.question_id for q in questions}, "逐题落库题数与输入不符"
@pytest.mark.asyncio
async def test_double_and_unit_total_and_orphan(self, harness_log: HarnessLog) -> None:
"""一次全链路推理同时锁定契约 3(双向 AND)、4(unit total)、5(孤儿剔除)。"""
questions = self._mixed_questions()
result, captured = await _run_capturing_warnings(questions, harness_log)
# 契约 43 single + 2 pair = 5 unitpair 计 1、孤儿不计入 total
assert result.total == 5, f"unit 粒度 total 错误:{result.total}"
# 契约 3single 3 对 + pairA 对 + pairB 错(双向 AND= 4
assert result.correct == 4, f"双向 AND 聚合错误:correct={result.correct}"
self._assert_task_type_breakdown(result)
# 契约 5:孤儿 pair 被告警(不静默)
assert any("orphan_p" in msg for msg in captured), "孤儿 pair 未告警(静默剔除)"
self._assert_all_persisted(harness_log, questions)
# ── 契约 6:纯非 AR build_batches 与旧逐题黄金参照逐字节一致 ────────────
class TestNonARByteIdentical:
"""纯 single 题库 build_batches 输出与引入 QuestionUnit 前的旧逐题逻辑逐字节一致。"""
def test_pure_single_matches_golden_reference(self) -> None:
"""混格链路对纯非 AR 输入不得引入任何漂移(黄金参照对照)。"""
items = [_single(f"q{i}", task_type=f"t{i % 4}") for i in range(40)]
correctness = {f"q{i}": (i % 3 == 0) for i in range(40)}
got, _ = build_batches(items, correctness, 8, 3, seed=7, correct_ratio=0.5)
ref = _reference_build_batches(items, correctness, 8, 3, seed=7, correct_ratio=0.5)
# 非空护栏:确保是实质性非空比较(防空==空误通过)
assert sum(len(b) for b in ref) > 0, "黄金参照为空,byte-identical 断言未实质覆盖"
assert _ids(got) == _ids(ref), "纯非 AR build_batches 与旧逐题黄金参照产生漂移"