Skip to content

Commit 50a6294

Browse files
committed
fix(examples): isolate fake model training evidence
1 parent 6e42d29 commit 50a6294

4 files changed

Lines changed: 281 additions & 76 deletions

File tree

examples/optimization/eval_optimize_loop/eval_loop/evaluator.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -29,7 +29,7 @@ def __init__(self, model: FakeModel, judge: FakeJudge, *, trace_enabled: bool =
2929
def evaluate(self, *, prompt_id: str, prompt: str, cases: Iterable[EvalCase], split: str) -> EvalResult:
3030
case_results: list[CaseResult] = []
3131
for case in cases:
32-
output, model_trace, cost = self.model.generate(prompt_id, prompt, case)
32+
output, model_trace, cost = self.model.generate(prompt_id, prompt, case.input)
3333
judged = self.judge.score(case, output)
3434
failure_category = None
3535
failure_reason = None

examples/optimization/eval_optimize_loop/eval_loop/fake_model.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -7,8 +7,6 @@
77
from dataclasses import dataclass
88
from typing import Any
99

10-
from .schemas import EvalCase
11-
1210
_ASSIGNMENT_PATTERN = re.compile(r"(?<![A-Za-z0-9_-])([A-Za-z_][A-Za-z0-9_]*)=([A-Za-z0-9_-]+)" r"(?![A-Za-z0-9_-])")
1311
_RETURN_ONLY_PATTERN = re.compile(
1412
r"\breturn\s+only\s+([A-Za-z0-9_-]+)\b",
@@ -40,10 +38,13 @@ def generate(
4038
self,
4139
prompt_id: str,
4240
prompt: str,
43-
case: EvalCase,
41+
user_input: str,
4442
) -> tuple[str, dict[str, Any], float]:
43+
if not isinstance(user_input, str):
44+
raise TypeError("user_input must be a string")
45+
4546
mode = self._mode(prompt)
46-
request = self._parse_request(case.input)
47+
request = self._parse_request(user_input)
4748
output = self._render(request, mode=mode)
4849
trace = {
4950
"seed": self.seed,

examples/optimization/eval_optimize_loop/eval_loop/optimizer.py

Lines changed: 54 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,8 @@
22

33
from __future__ import annotations
44

5+
from collections import Counter
6+
57
from .diffing import make_unified_diff
68
from .schemas import CandidatePrompt
79
from .schemas import EvalResult
@@ -23,18 +25,30 @@ def propose(
2325
baseline_train: EvalResult,
2426
failure_summary: dict[str, object],
2527
) -> list[CandidatePrompt]:
26-
failed_cases = [case for case in baseline_train.cases if not case.passed]
27-
if not failed_cases:
28-
return []
28+
if not isinstance(failure_summary, dict):
29+
raise TypeError("failure_summary must be a dict")
2930

30-
observed_categories = {case.failure_category for case in failed_cases if case.failure_category}
31-
by_category = failure_summary.get("by_category")
32-
if isinstance(by_category, dict):
33-
observed_categories.update(
34-
str(category) for category, count in by_category.items() if _is_positive_count(count)
31+
if baseline_train.split != "train":
32+
raise ValueError(
33+
"baseline_train.split must be 'train'; "
34+
f"got {baseline_train.split!r}"
3535
)
36+
for case in baseline_train.cases:
37+
if case.split != "train":
38+
raise ValueError(
39+
f"baseline_train case {case.case_id!r} split must be 'train'; "
40+
f"got {case.split!r}"
41+
)
3642

37-
targeted = sorted(observed_categories & _TARGET_FAILURE_CATEGORIES)
43+
failed_cases = [case for case in baseline_train.cases if not case.passed]
44+
observed_counts = Counter(
45+
case.failure_category
46+
for case in failed_cases
47+
if case.failure_category
48+
)
49+
_validate_failure_summary(failure_summary, observed_counts)
50+
51+
targeted = sorted(observed_counts.keys() & _TARGET_FAILURE_CATEGORIES)
3852
if not targeted:
3953
return []
4054

@@ -75,5 +89,34 @@ def propose(
7589
]
7690

7791

78-
def _is_positive_count(value: object) -> bool:
79-
return not isinstance(value, bool) and isinstance(value, (int, float)) and value > 0
92+
def _validate_failure_summary(
93+
failure_summary: dict[str, object],
94+
observed_counts: Counter[str],
95+
) -> None:
96+
if "by_category" not in failure_summary:
97+
return
98+
99+
by_category = failure_summary["by_category"]
100+
if not isinstance(by_category, dict):
101+
raise ValueError("failure_summary['by_category'] must be a dict")
102+
103+
summary_counts: dict[object, int] = {}
104+
for category, count in by_category.items():
105+
summary_counts[category] = _normalize_positive_count(category, count)
106+
107+
if summary_counts != dict(observed_counts):
108+
raise ValueError(
109+
"failure_summary['by_category'] must match failed train cases exactly; "
110+
f"summary={summary_counts!r}, observed={dict(observed_counts)!r}"
111+
)
112+
113+
114+
def _normalize_positive_count(category: object, value: object) -> int:
115+
normalized = value if type(value) is int and value > 0 else None
116+
117+
if normalized is None:
118+
raise ValueError(
119+
"failure_summary['by_category'] count must be a positive integer; "
120+
f"category={category!r}, count={value!r}"
121+
)
122+
return normalized

0 commit comments

Comments
 (0)