Skip to content

Commit 4563b0c

Browse files
committed
refactor(examples): unify evaluation optimization pipeline
1 parent 50a6294 commit 4563b0c

8 files changed

Lines changed: 2169 additions & 801 deletions

File tree

examples/optimization/eval_optimize_loop/eval_loop/gate.py

Lines changed: 109 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55
from typing import Any
66

77
from .schemas import CaseDelta
8+
from .schemas import CostSummary
89
from .schemas import EvalResult
910
from .schemas import GateDecision
1011

@@ -35,17 +36,32 @@ def decide(
3536
candidate_train: EvalResult,
3637
candidate_validation: EvalResult,
3738
deltas: list[CaseDelta],
39+
cost_summary: CostSummary,
3840
cumulative_cost: float = 0.0,
3941
) -> GateDecision:
4042
train_delta = round(candidate_train.score - baseline_train.score, 6)
4143
val_delta = round(candidate_validation.score - baseline_validation.score, 6)
4244
candidate_cost = round(candidate_train.cost + candidate_validation.cost, 6)
4345
reasons: list[str] = []
4446

47+
train_comparable = _append_comparability_reasons(
48+
reasons,
49+
split="train",
50+
baseline=baseline_train,
51+
candidate=candidate_train,
52+
)
53+
validation_comparable = _append_comparability_reasons(
54+
reasons,
55+
split="validation",
56+
baseline=baseline_validation,
57+
candidate=candidate_validation,
58+
)
59+
4560
overfit_detected = train_delta > 0 and val_delta <= 0
4661
if overfit_detected:
4762
reasons.append(
48-
"reject: overfit detected because train score improved but validation score regressed or did not improve "
63+
"reject: overfit detected because train score improved but "
64+
"validation score regressed or did not improve "
4965
f"({train_delta:+.3f} train, {val_delta:+.3f} validation)"
5066
)
5167

@@ -56,20 +72,32 @@ def decide(
5672
f"{val_delta:+.3f} is below required {min_val_improvement:+.3f}"
5773
)
5874

59-
baseline_validation_by_id = baseline_validation.by_case_id()
60-
candidate_validation_by_id = candidate_validation.by_case_id()
61-
validation_new_failures = [
62-
case_id
63-
for case_id, candidate_case in sorted(candidate_validation_by_id.items())
64-
if not candidate_case.passed and baseline_validation_by_id.get(case_id)
65-
and baseline_validation_by_id[case_id].passed
66-
]
67-
new_hard_failures = [
68-
case_id
69-
for case_id, candidate_case in sorted(candidate_validation_by_id.items())
70-
if candidate_case.hard_failed and baseline_validation_by_id.get(case_id)
71-
and baseline_validation_by_id[case_id].passed
72-
]
75+
validation_new_failures: list[str] = []
76+
if validation_comparable:
77+
baseline_validation_by_id = baseline_validation.by_case_id()
78+
candidate_validation_by_id = candidate_validation.by_case_id()
79+
validation_new_failures = [
80+
case_id
81+
for case_id, candidate_case in sorted(candidate_validation_by_id.items())
82+
if not candidate_case.passed and baseline_validation_by_id[case_id].passed
83+
]
84+
if validation_new_failures:
85+
reasons.append(f"reject: new validation failures appeared: {validation_new_failures}")
86+
87+
new_hard_failures: list[str] = []
88+
for baseline_result, candidate_result, comparable in (
89+
(baseline_train, candidate_train, train_comparable),
90+
(baseline_validation, candidate_validation, validation_comparable),
91+
):
92+
if not comparable:
93+
continue
94+
baseline_by_id = baseline_result.by_case_id()
95+
new_hard_failures.extend(
96+
case_id
97+
for case_id, candidate_case in sorted(candidate_result.by_case_id().items())
98+
if candidate_case.hard_failed and not baseline_by_id[case_id].hard_failed
99+
)
100+
new_hard_failures = sorted(set(new_hard_failures))
73101
if new_hard_failures and not bool(self.config["allow_new_hard_fail"]):
74102
reasons.append(f"reject: new hard failures appeared: {new_hard_failures}")
75103

@@ -91,10 +119,16 @@ def decide(
91119
if excessive_drops:
92120
reasons.append(f"reject: per-case validation score drops exceed {max_drop:.3f}: {excessive_drops}")
93121

94-
max_total_cost = float(self.config["max_total_cost"])
95122
total_run_cost = round(cumulative_cost + candidate_cost, 6)
96-
if total_run_cost > max_total_cost:
97-
reasons.append(f"reject: total run cost {total_run_cost:.3f} exceeds budget {max_total_cost:.3f}")
123+
configured_max_total_cost = self.config["max_total_cost"]
124+
if configured_max_total_cost is not None:
125+
max_total_cost = float(configured_max_total_cost)
126+
if not cost_summary.complete:
127+
reasons.append("reject: cost_unavailable for configured max_total_cost")
128+
elif total_run_cost > max_total_cost:
129+
reasons.append(
130+
f"reject: total run cost {total_run_cost:.3f} exceeds budget {max_total_cost:.3f}"
131+
)
98132

99133
accepted = not any(reason.startswith("reject:") for reason in reasons)
100134
if accepted:
@@ -118,4 +152,61 @@ def decide(
118152
cumulative_cost=round(cumulative_cost, 6),
119153
total_run_cost=total_run_cost,
120154
cost=candidate_cost,
155+
gate_status="applied",
156+
gate_not_applied_reason=None,
157+
not_applied_checks=[],
158+
)
159+
160+
161+
def _append_comparability_reasons(
162+
reasons: list[str],
163+
*,
164+
split: str,
165+
baseline: EvalResult,
166+
candidate: EvalResult,
167+
) -> bool:
168+
"""Reject malformed result pairs before any dict conversion can hide evidence."""
169+
170+
comparable = True
171+
baseline_ids = [case.case_id for case in baseline.cases]
172+
candidate_ids = [case.case_id for case in candidate.cases]
173+
if not baseline_ids:
174+
reasons.append(f"reject: baseline {split} has empty case results")
175+
comparable = False
176+
if not candidate_ids:
177+
reasons.append(f"reject: candidate {split} has empty case results")
178+
comparable = False
179+
180+
baseline_duplicates = _duplicates(baseline_ids)
181+
candidate_duplicates = _duplicates(candidate_ids)
182+
if baseline_duplicates:
183+
reasons.append(
184+
f"reject: baseline {split} has duplicate case IDs: {baseline_duplicates}"
185+
)
186+
comparable = False
187+
if candidate_duplicates:
188+
reasons.append(
189+
f"reject: candidate {split} has duplicate case IDs: {candidate_duplicates}"
190+
)
191+
comparable = False
192+
193+
baseline_set = set(baseline_ids)
194+
candidate_set = set(candidate_ids)
195+
if baseline_set != candidate_set:
196+
reasons.append(
197+
f"reject: {split} case ID set mismatch; "
198+
f"missing={sorted(baseline_set - candidate_set)}, "
199+
f"extra={sorted(candidate_set - baseline_set)}"
121200
)
201+
comparable = False
202+
return comparable
203+
204+
205+
def _duplicates(case_ids: list[str]) -> list[str]:
206+
seen: set[str] = set()
207+
duplicates: set[str] = set()
208+
for case_id in case_ids:
209+
if case_id in seen:
210+
duplicates.add(case_id)
211+
seen.add(case_id)
212+
return sorted(duplicates)

0 commit comments

Comments
 (0)