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
+43 -1
View File
@@ -120,6 +120,9 @@ class _TrainState:
gate_epoch_observed: bool = False gate_epoch_observed: bool = False
probations: dict[str, Probation] = field(default_factory=dict) probations: dict[str, Probation] = field(default_factory=dict)
gate_cooldown: dict[str, int] = field(default_factory=dict) gate_cooldown: dict[str, int] = field(default_factory=dict)
# 进程内 holdout 去重备忘录 (skills_v, prompts_v) -> test 评估结果;不进 checkpoint
# resume 后清空(重评一次是可接受代价,换取零 schema 变更)。
holdout_memo: dict[tuple[str, str], InferenceResult] = field(default_factory=dict)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -2006,11 +2009,18 @@ class Runner:
) -> None: ) -> None:
"""四向 held-outbaseline/best_hard/final/best_mixed 各在 test 池评估。 """四向 held-outbaseline/best_hard/final/best_mixed 各在 test 池评估。
去重(进程内备忘录 state.holdout_memo,不改 schema):
- baseline:不跑推理,从基线 predictions 推导 test 结果(0 推理),存 memo 跨 epoch 复用。
- final:真评 test,存 memo[(final_sv,final_pv)]。
- best_hard:其版本已在 memo== final 或往轮已评)则引用,否则真评并存 memo。
- best_mixed:赢家必是 best_hard 或 final 之一,其结果已在 memo,直接引用(0 推理)。
test 池仅观测落库,绝不进 gate/best/early-stop/调参。 test 池仅观测落库,绝不进 gate/best/early-stop/调参。
""" """
best_mixed = await self._pick_mixed_best( best_mixed = await self._pick_mixed_best(
epoch, pools, state, eval_skills_version, eval_prompts_version epoch, pools, state, eval_skills_version, eval_prompts_version
) )
memo = state.holdout_memo
versions: dict[str, tuple[str, str] | None] = { versions: dict[str, tuple[str, str] | None] = {
"baseline": (state.baseline_skills_version, state.baseline_prompts_version), "baseline": (state.baseline_skills_version, state.baseline_prompts_version),
"best_hard": (state.best_skills_version, state.best_prompts_version), "best_hard": (state.best_skills_version, state.best_prompts_version),
@@ -2022,9 +2032,14 @@ class Runner:
continue continue
sv, pv = version sv, pv = version
run_id = f"{self._config.run_id}_holdout_{version_kind}_e{epoch}" run_id = f"{self._config.run_id}_holdout_{version_kind}_e{epoch}"
res = await self._eval_version_on_pool( if version not in memo:
if version_kind == "baseline":
memo[version] = self._derive_baseline_test_result(pools, run_id)
else:
memo[version] = await self._eval_version_on_pool(
sv, pv, pools.test, run_id, context=f"held-out {version_kind}" sv, pv, pools.test, run_id, context=f"held-out {version_kind}"
) )
res = memo[version]
soft = await self._try_soft_score(run_id, pools.test) soft = await self._try_soft_score(run_id, pools.test)
mixed = None if soft is None else 0.5 * res.accuracy + 0.5 * soft mixed = None if soft is None else 0.5 * res.accuracy + 0.5 * soft
write_holdout_eval( write_holdout_eval(
@@ -2038,6 +2053,33 @@ class Runner:
per_task_type_json=json.dumps(res.per_task_type, ensure_ascii=False), per_task_type_json=json.dumps(res.per_task_type, ensure_ascii=False),
) )
def _derive_baseline_test_result(self, pools: Pools, run_id: str) -> InferenceResult:
"""从基线 run 的 predictions 推导 test 池评估结果(0 推理)。
基线 runpools.baseline_run_id)已对全题库推理并落库,test 题在其中;此处
按 test 题回读基线预测、经 _aggregate_results 折叠为 unit 级 InferenceResult
避免重复推理基线版本(同版本不重采样)。
参数:
pools: 冻结三池(提供 test 与 baseline_run_id)。
run_id: 本次 holdout baseline 向的 run_id(仅用作结果标识)。
返回:
unit 级 InferenceResultaccuracy / per_task_type 与真评同口径)。
"""
from app.harness.inference import _aggregate_results
from app.harness.log import HarnessLog
qids = [q.question_id for q in pools.test]
with HarnessLog(
str(self._paths.db_path), pools.baseline_run_id, register_run=False
) as log:
placeholders = ", ".join(["?"] * len(qids))
rows = log.query(
f"SELECT * FROM predictions WHERE run_id=? AND question_id IN ({placeholders})",
(pools.baseline_run_id, *qids),
)
return _aggregate_results(rows, pools.test, run_id)
async def _pick_mixed_best( async def _pick_mixed_best(
self, self,
epoch: int, epoch: int,
@@ -212,6 +212,63 @@ async def test_diagnosis_reads_traces_from_steps_json(
assert traces[0]["tool_name"] == "search_tree" 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 @pytest.mark.asyncio
async def test_run_step_deletes_stale_rows_before_rerun( async def test_run_step_deletes_stale_rows_before_rerun(
runner_with_real_store: Runner, monkeypatch: pytest.MonkeyPatch runner_with_real_store: Runner, monkeypatch: pytest.MonkeyPatch