perf: dedup holdout four-way eval (baseline derive, best_hard memo)

This commit is contained in:
2026-07-16 06:31:00 -04:00
parent 77fd35830c
commit 58a0203522
2 changed files with 102 additions and 3 deletions
@@ -212,6 +212,63 @@ async def test_diagnosis_reads_traces_from_steps_json(
assert traces[0]["tool_name"] == "search_tree"
def _fake_inference_result(accuracy: float) -> object:
"""构造 InferenceResult 供 holdout 去重测试(per_task_type 空、token 归零)。"""
from app.harness.inference import InferenceResult
return InferenceResult(
run_id="x",
accuracy=accuracy,
total=10,
correct=int(accuracy * 10),
per_task_type={},
steps_mean=1.0,
token_usage={"prompt_tokens": 0, "completion_tokens": 0},
stop_reason_counts={},
)
@pytest.mark.asyncio
async def test_holdout_dedup_skips_reevaluated_versions(
runner_with_real_store: Runner,
) -> None:
"""baseline 不跑推理(从基线预测推导);best_hard==final 时不重复评估。"""
from unittest.mock import AsyncMock
runner = runner_with_real_store
eval_calls: list[str] = []
async def _fake_eval(sv, pv, questions, run_id, context): # noqa: ANN001
eval_calls.append(run_id)
return _fake_inference_result(0.5)
runner._eval_version_on_pool = AsyncMock(side_effect=_fake_eval)
# best_mixed 赢家取 final 版本(必落在 memo,0 推理)
runner._pick_mixed_best = AsyncMock(return_value=("skills_final", "prompts_final"))
# baseline 推导:patch 为 0 推理的假结果(不经 _eval_version_on_pool
runner._derive_baseline_test_result = MagicMock(return_value=_fake_inference_result(0.4))
state = MagicMock()
state.holdout_memo = {}
state.baseline_skills_version = "skills_base"
state.baseline_prompts_version = "prompts_base"
# best_hard 版本 == final 版本 → 去重
state.best_skills_version = "skills_final"
state.best_prompts_version = "prompts_final"
pools = MagicMock()
pools.baseline_run_id = "infer_adhoc"
pools.test = [_fake_question("q1", "vA")]
await runner._holdout_four_way(
1, pools, state, eval_skills_version="skills_final", eval_prompts_version="prompts_final"
)
# baseline=0(推导)、best_hard=1(真评并存 memo)、final/best_mixed 引用 memo0
assert len(eval_calls) == 1
runner._derive_baseline_test_result.assert_called_once()
@pytest.mark.asyncio
async def test_run_step_deletes_stale_rows_before_rerun(
runner_with_real_store: Runner, monkeypatch: pytest.MonkeyPatch