downshift 0.1.0.dev0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
downshift/decide.py ADDED
@@ -0,0 +1,220 @@
1
+ """Downgrade decisions: the cheapest model that keeps enough of the baseline's quality.
2
+
3
+ Decisions use pass rate (share of cases fully right), never mean score. A model can get
4
+ most JSON fields right and still fail most cases.
5
+
6
+ A candidate qualifies when all of these hold:
7
+ - its results are complete (every case scored, no error rows),
8
+ - pass rate / baseline pass rate >= threshold,
9
+ - pass rate >= min_pass_rate (absolute floor, 0 = off),
10
+ - it costs less per call than the baseline (priced with its own measured tokens).
11
+ The cheapest qualifying candidate wins; ties go to the higher pass rate, then list order.
12
+ """
13
+
14
+ from __future__ import annotations
15
+
16
+ from collections.abc import Iterable, Mapping, Sequence
17
+ from dataclasses import dataclass
18
+ from pathlib import Path
19
+
20
+ from downshift.config import ModelPrice
21
+ from downshift.evals import EvalSet
22
+ from downshift.runner import ResultRow, RunSummary, load_results, results_path, summarize_rows
23
+
24
+ KEEP = "keep"
25
+ DOWNGRADE = "downgrade"
26
+ EPSILON = 1e-9 # 18/22 vs 20/22 must count as exactly 0.90
27
+
28
+
29
+ def call_cost(price: ModelPrice, prompt_tokens: float, completion_tokens: float) -> float:
30
+ """USD for one call with these (average) token counts."""
31
+ total = prompt_tokens * price.input_per_mtok + completion_tokens * price.output_per_mtok
32
+ return total / 1_000_000
33
+
34
+
35
+ @dataclass(frozen=True)
36
+ class ModelStats:
37
+ """Results of one model on one call site."""
38
+
39
+ summary: RunSummary
40
+ avg_judge_score: float | None = None
41
+
42
+ @property
43
+ def model(self) -> str:
44
+ return self.summary.model
45
+
46
+ @property
47
+ def pass_rate(self) -> float | None:
48
+ return self.summary.pass_rate
49
+
50
+ @property
51
+ def complete(self) -> bool:
52
+ s = self.summary
53
+ return s.cases > 0 and s.errors == 0 and s.scored == s.cases
54
+
55
+ def cost_per_call(self, price: ModelPrice) -> float | None:
56
+ s = self.summary
57
+ if s.avg_prompt_tokens is None or s.avg_completion_tokens is None:
58
+ return None
59
+ return call_cost(price, s.avg_prompt_tokens, s.avg_completion_tokens)
60
+
61
+
62
+ def stats_for(
63
+ site_id: str, model: str, case_ids: Sequence[str], rows: Mapping[str, ResultRow]
64
+ ) -> ModelStats:
65
+ """Stats for these case ids. Avg judge score counts only scored rows that have one."""
66
+ summary = summarize_rows(site_id, model, case_ids, rows)
67
+ judged = [
68
+ float(rows[c].judge_score) # type: ignore[arg-type]
69
+ for c in case_ids
70
+ if c in rows and rows[c].ok and rows[c].judge_score is not None
71
+ ]
72
+ avg = sum(judged) / len(judged) if judged else None
73
+ return ModelStats(summary=summary, avg_judge_score=avg)
74
+
75
+
76
+ def load_site_stats(
77
+ site_id: str, eval_set: EvalSet, models: Iterable[str], results_dir: Path
78
+ ) -> dict[str, ModelStats]:
79
+ """Stats per model for one site, from results files (last row per case wins)."""
80
+ case_ids = [case.id for case in eval_set.cases]
81
+ out: dict[str, ModelStats] = {}
82
+ for model in models:
83
+ rows = load_results(results_path(results_dir, site_id, model))
84
+ out[model] = stats_for(site_id, model, case_ids, rows)
85
+ return out
86
+
87
+
88
+ @dataclass(frozen=True)
89
+ class CandidateCheck:
90
+ """Why one candidate did or did not qualify."""
91
+
92
+ model: str
93
+ pass_rate: float | None
94
+ ratio: float | None
95
+ cost_per_call: float | None
96
+ passed: bool
97
+ reason: str
98
+
99
+
100
+ @dataclass(frozen=True)
101
+ class Decision:
102
+ site_id: str
103
+ action: str # KEEP or DOWNGRADE
104
+ model: str # the model to use after the decision
105
+ baseline: str
106
+ reason: str
107
+ baseline_pass_rate: float | None
108
+ baseline_below_floor: bool
109
+ checks: tuple[CandidateCheck, ...]
110
+ stats: Mapping[str, ModelStats]
111
+
112
+ @property
113
+ def downgraded(self) -> bool:
114
+ return self.action == DOWNGRADE
115
+
116
+
117
+ def _price(prices: Mapping[str, ModelPrice], model: str) -> ModelPrice:
118
+ try:
119
+ return prices[model]
120
+ except KeyError:
121
+ raise ValueError(f"no pricing for model {model!r}") from None
122
+
123
+
124
+ def _incomplete(s: RunSummary) -> str:
125
+ return f"incomplete results ({s.errors} errors, {s.scored}/{s.cases} scored)"
126
+
127
+
128
+ def _check(
129
+ model: str,
130
+ s: ModelStats | None,
131
+ *,
132
+ base_rate: float,
133
+ base_cost: float | None,
134
+ prices: Mapping[str, ModelPrice],
135
+ threshold: float,
136
+ min_pass_rate: float,
137
+ ) -> CandidateCheck:
138
+ if s is None or s.summary.scored == 0:
139
+ return CandidateCheck(model, None, None, None, False, "no results")
140
+ rate = s.pass_rate or 0.0
141
+ ratio = rate / base_rate
142
+ cost = s.cost_per_call(_price(prices, model))
143
+
144
+ def fail(reason: str) -> CandidateCheck:
145
+ return CandidateCheck(model, rate, ratio, cost, False, reason)
146
+
147
+ if not s.complete:
148
+ return fail(_incomplete(s.summary))
149
+ if ratio < threshold - EPSILON:
150
+ return fail(f"keeps {ratio:.0%} of baseline quality, needs {threshold:.0%}")
151
+ if rate < min_pass_rate - EPSILON:
152
+ return fail(f"pass rate {rate:.0%} is below the floor {min_pass_rate:.0%}")
153
+ if cost is None or base_cost is None or cost >= base_cost:
154
+ return fail("not cheaper than the baseline")
155
+ return CandidateCheck(
156
+ model, rate, ratio, cost, True, f"keeps {ratio:.0%} of baseline quality and costs less"
157
+ )
158
+
159
+
160
+ def decide_site(
161
+ site_id: str,
162
+ stats: Mapping[str, ModelStats],
163
+ *,
164
+ baseline: str,
165
+ candidates: Sequence[str],
166
+ prices: Mapping[str, ModelPrice],
167
+ threshold: float,
168
+ min_pass_rate: float = 0.0,
169
+ ) -> Decision:
170
+ """Keep the baseline or downgrade to the cheapest qualifying candidate."""
171
+ if not 0.0 < threshold <= 1.0:
172
+ raise ValueError(f"threshold must be in (0, 1], got {threshold}")
173
+ if not 0.0 <= min_pass_rate <= 1.0:
174
+ raise ValueError(f"min_pass_rate must be in [0, 1], got {min_pass_rate}")
175
+
176
+ def keep(
177
+ reason: str,
178
+ base_rate: float | None = None,
179
+ below: bool = False,
180
+ checks: tuple[CandidateCheck, ...] = (),
181
+ ) -> Decision:
182
+ return Decision(site_id, KEEP, baseline, baseline, reason, base_rate, below, checks, stats)
183
+
184
+ base = stats.get(baseline)
185
+ if base is None or base.summary.scored == 0:
186
+ return keep("no baseline results")
187
+ if not base.complete:
188
+ return keep(f"baseline has {_incomplete(base.summary)}", base.pass_rate)
189
+ base_rate = base.pass_rate or 0.0
190
+ below = base_rate < min_pass_rate - EPSILON
191
+ if base_rate == 0.0:
192
+ return keep("baseline passes no cases; nothing to compare against", 0.0, below)
193
+
194
+ base_cost = base.cost_per_call(_price(prices, baseline))
195
+ checks = tuple(
196
+ _check(
197
+ model,
198
+ stats.get(model),
199
+ base_rate=base_rate,
200
+ base_cost=base_cost,
201
+ prices=prices,
202
+ threshold=threshold,
203
+ min_pass_rate=min_pass_rate,
204
+ )
205
+ for model in candidates
206
+ if model != baseline
207
+ )
208
+ order = {c.model: i for i, c in enumerate(checks)}
209
+ passing = [c for c in checks if c.passed]
210
+ if not passing:
211
+ reason = "no candidate models" if not checks else "no cheaper model passed the checks"
212
+ return keep(reason, base_rate, below, checks)
213
+
214
+ best = min(
215
+ passing,
216
+ key=lambda c: (c.cost_per_call or 0.0, -(c.pass_rate or 0.0), order[c.model]),
217
+ )
218
+ return Decision(
219
+ site_id, DOWNGRADE, best.model, baseline, best.reason, base_rate, below, checks, stats
220
+ )
downshift/evalgen.py ADDED
@@ -0,0 +1,220 @@
1
+ """Generate eval sets with any configured model (the no-Bob path).
2
+
3
+ `generate_eval_set` asks a model for test cases for one call site, then keeps
4
+ only the cases that pass the same checks as `downshift check-evals`. Cases are
5
+ renumbered, so the model never has to get ids right. The CLI command
6
+ `evalgen` is a thin wrapper that writes one `<slug>.jsonl` per call site.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import json
12
+ import os
13
+ import re
14
+ from dataclasses import dataclass, field
15
+ from pathlib import Path
16
+ from typing import Any
17
+
18
+ from downshift.evals import (
19
+ EvalCase,
20
+ EvalSet,
21
+ allowed_values,
22
+ expression_placeholders,
23
+ graded_fields,
24
+ site_placeholders,
25
+ validate_eval_set,
26
+ )
27
+ from downshift.llm import LLMClient
28
+ from downshift.schema import CallSite
29
+
30
+ _FENCE = re.compile(r"^```(?:json)?\s*|\s*```$")
31
+
32
+ SYSTEM_PROMPT = (
33
+ "You write test cases for one LLM feature in a software product. "
34
+ "Every expected answer must be certainly correct. Reply with JSON only."
35
+ )
36
+
37
+
38
+ class EvalGenSkip(Exception):
39
+ """Raised when a call site cannot get generated evals; the message says why."""
40
+
41
+
42
+ @dataclass
43
+ class GenerationResult:
44
+ site_id: str
45
+ grading: str
46
+ cases: list[EvalCase] = field(default_factory=list)
47
+ dropped: list[str] = field(default_factory=list)
48
+ attempts: int = 0
49
+
50
+
51
+ def grading_for(site: CallSite) -> str:
52
+ """The audit's grading, or a safe default for an unaudited call site."""
53
+ if site.grading:
54
+ return site.grading
55
+ return "json_fields" if site.output_format == "json" else "judge"
56
+
57
+
58
+ def generate_eval_set(
59
+ client: LLMClient,
60
+ model: str,
61
+ site: CallSite,
62
+ *,
63
+ count: int = 20,
64
+ shared: dict[str, str] | None = None,
65
+ max_attempts: int = 3,
66
+ temperature: float = 0.7,
67
+ ) -> GenerationResult:
68
+ """Ask `model` for cases until `count` valid ones are kept or attempts run out."""
69
+ _check_site(site)
70
+ shared = {k: v for k, v in (shared or {}).items() if k in site_placeholders(site)}
71
+ grading = grading_for(site)
72
+ result = GenerationResult(site_id=site.id, grading=grading)
73
+ seen_inputs: set[str] = set()
74
+
75
+ while len(result.cases) < count and result.attempts < max_attempts:
76
+ result.attempts += 1
77
+ wanted = count - len(result.cases)
78
+ messages = build_messages(site, grading, wanted, shared)
79
+ completion = client.complete(model, messages, temperature=temperature, json_mode=True)
80
+ try:
81
+ raw_cases = parse_cases(completion.text)
82
+ except ValueError as exc:
83
+ result.dropped.append(f"attempt {result.attempts}: {exc}")
84
+ continue
85
+ for raw in raw_cases:
86
+ if len(result.cases) >= count:
87
+ break
88
+ case, problem = _to_case(raw, site, grading, shared, len(result.cases) + 1)
89
+ if case is None:
90
+ result.dropped.append(problem)
91
+ continue
92
+ key = json.dumps(case.inputs, sort_keys=True, ensure_ascii=False)
93
+ if key in seen_inputs:
94
+ result.dropped.append(f"duplicate inputs: {case.notes}")
95
+ continue
96
+ seen_inputs.add(key)
97
+ result.cases.append(case)
98
+ return result
99
+
100
+
101
+ def build_messages(
102
+ site: CallSite, grading: str, count: int, shared: dict[str, str]
103
+ ) -> list[dict[str, str]]:
104
+ inputs = [name for name in site_placeholders(site) if name not in shared]
105
+ prompt_lines = [f"[{m.role}]\n{m.content}" for m in site.messages or []]
106
+ parts = [
107
+ f"Feature: {site.purpose or site.function}",
108
+ f"Output contract: {site.output_contract or 'not documented'}",
109
+ "Prompt template sent to the model ({name} is filled per case):",
110
+ "\n\n".join(prompt_lines),
111
+ ]
112
+ if shared:
113
+ parts.append("Fixed inputs, the same for every case (do not repeat them in cases):")
114
+ parts.extend(f"--- {k} ---\n{v}" for k, v in shared.items())
115
+ parts.append(f"Write {count} diverse, realistic test cases.")
116
+ parts.append(
117
+ f"Each case has keys: inputs (an object with exactly these keys: "
118
+ f"{', '.join(inputs)}), expected, notes (what the case tests)."
119
+ )
120
+ parts.append(_expected_rule(site, grading))
121
+ parts.append(
122
+ "Include edge cases: very short text, other languages, typos, "
123
+ "missing details. Skip any case whose answer could be argued."
124
+ )
125
+ parts.append('Reply with: {"cases": [ ... ]}')
126
+ return [
127
+ {"role": "system", "content": SYSTEM_PROMPT},
128
+ {"role": "user", "content": "\n\n".join(parts)},
129
+ ]
130
+
131
+
132
+ def parse_cases(text: str) -> list[Any]:
133
+ """Pull the case list out of a model reply. Raises ValueError if there is none."""
134
+ cleaned = _FENCE.sub("", text.strip())
135
+ try:
136
+ data = json.loads(cleaned)
137
+ except json.JSONDecodeError as exc:
138
+ raise ValueError(f"reply is not valid JSON ({exc.msg})") from exc
139
+ if isinstance(data, dict):
140
+ data = data.get("cases")
141
+ if not isinstance(data, list):
142
+ raise ValueError('reply has no "cases" list')
143
+ return data
144
+
145
+
146
+ def write_eval_set(
147
+ path: Path, cases: list[EvalCase], shared_files: dict[str, Path] | None = None
148
+ ) -> None:
149
+ """Write cases as JSONL, with a `shared` first line pointing at files if given."""
150
+ path.parent.mkdir(parents=True, exist_ok=True)
151
+ lines = []
152
+ if shared_files:
153
+ header = {
154
+ key: {"file": Path(os.path.relpath(target, path.parent)).as_posix()}
155
+ for key, target in shared_files.items()
156
+ }
157
+ lines.append(json.dumps({"shared": header}, ensure_ascii=False))
158
+ for case in cases:
159
+ record = {
160
+ "id": case.id,
161
+ "inputs": case.inputs,
162
+ "expected": case.expected,
163
+ "grading": case.grading,
164
+ "notes": case.notes,
165
+ }
166
+ lines.append(json.dumps(record, ensure_ascii=False))
167
+ path.write_text("\n".join(lines) + "\n", encoding="utf-8")
168
+
169
+
170
+ # --- helpers ------------------------------------------------------------------
171
+
172
+
173
+ def _check_site(site: CallSite) -> None:
174
+ if not site.prompt_resolved:
175
+ raise EvalGenSkip("prompt is not resolved; run the Bob auditor first")
176
+ expressions = expression_placeholders(site)
177
+ if expressions:
178
+ names = ", ".join("{" + e + "}" for e in expressions)
179
+ raise EvalGenSkip(f"prompt has expression placeholders {names}; run the Bob auditor first")
180
+
181
+
182
+ def _expected_rule(site: CallSite, grading: str) -> str:
183
+ if grading == "exact":
184
+ labels = allowed_values(site)
185
+ if labels:
186
+ return f"expected: exactly one of {', '.join(labels)} (a string)."
187
+ return "expected: the exact short answer as a string."
188
+ if grading == "json_fields":
189
+ fields = graded_fields(site)
190
+ if fields:
191
+ return f"expected: an object with exactly these keys: {', '.join(fields)}."
192
+ return "expected: the JSON object the feature should return."
193
+ return (
194
+ "expected: a rubric string listing what a correct reply must and must not do, "
195
+ "point by point."
196
+ )
197
+
198
+
199
+ def _to_case(
200
+ raw: Any, site: CallSite, grading: str, shared: dict[str, str], number: int
201
+ ) -> tuple[EvalCase | None, str]:
202
+ if not isinstance(raw, dict):
203
+ return None, f"not an object: {str(raw)[:60]}"
204
+ notes = raw.get("notes")
205
+ notes = notes if isinstance(notes, str) and notes.strip() else "generated"
206
+ inputs = raw.get("inputs")
207
+ if not isinstance(inputs, dict):
208
+ return None, f"inputs is not an object: {notes}"
209
+ case = EvalCase(
210
+ id=f"{site.function}-{number:02d}",
211
+ inputs=inputs,
212
+ expected=raw.get("expected"),
213
+ grading=grading,
214
+ notes=notes,
215
+ )
216
+ check = validate_eval_set(EvalSet(path=Path("-"), cases=[case], shared=shared), site)
217
+ if check.errors:
218
+ reason = check.errors[0].split(": ", 1)[-1]
219
+ return None, f"{reason}: {notes}"
220
+ return case, ""