quantdiff 0.1.0rc1__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.
- quantdiff/__init__.py +53 -0
- quantdiff/__main__.py +5 -0
- quantdiff/_http.py +151 -0
- quantdiff/_text.py +13 -0
- quantdiff/_version.py +1 -0
- quantdiff/api.py +340 -0
- quantdiff/backends/__init__.py +28 -0
- quantdiff/backends/_common.py +342 -0
- quantdiff/backends/base.py +91 -0
- quantdiff/backends/llamacpp.py +428 -0
- quantdiff/backends/ollama.py +359 -0
- quantdiff/backends/openai_compat.py +338 -0
- quantdiff/cache.py +240 -0
- quantdiff/card.py +1664 -0
- quantdiff/cli.py +377 -0
- quantdiff/discover.py +488 -0
- quantdiff/errors.py +45 -0
- quantdiff/metrics/__init__.py +36 -0
- quantdiff/metrics/codeexec.py +428 -0
- quantdiff/metrics/jsonschema.py +610 -0
- quantdiff/metrics/logit.py +214 -0
- quantdiff/metrics/tasks.py +114 -0
- quantdiff/metrics/textsim.py +66 -0
- quantdiff/metrics/toolcheck.py +99 -0
- quantdiff/png.py +360 -0
- quantdiff/preflight.py +365 -0
- quantdiff/progress.py +283 -0
- quantdiff/py.typed +0 -0
- quantdiff/report.py +780 -0
- quantdiff/runner.py +492 -0
- quantdiff/spec.py +154 -0
- quantdiff/stats.py +226 -0
- quantdiff/suites/__init__.py +462 -0
- quantdiff/suites/data/chat.jsonl +22 -0
- quantdiff/suites/data/code.jsonl +32 -0
- quantdiff/suites/data/json.jsonl +34 -0
- quantdiff/suites/data/scoring.jsonl +41 -0
- quantdiff/suites/data/tools.jsonl +32 -0
- quantdiff/types.py +322 -0
- quantdiff/verdict.py +1513 -0
- quantdiff-0.1.0rc1.dist-info/METADATA +514 -0
- quantdiff-0.1.0rc1.dist-info/RECORD +45 -0
- quantdiff-0.1.0rc1.dist-info/WHEEL +4 -0
- quantdiff-0.1.0rc1.dist-info/entry_points.txt +2 -0
- quantdiff-0.1.0rc1.dist-info/licenses/LICENSE +202 -0
quantdiff/stats.py
ADDED
|
@@ -0,0 +1,226 @@
|
|
|
1
|
+
"""Paired significance tests and confidence intervals for comparing two models.
|
|
2
|
+
|
|
3
|
+
Every candidate answers the same cases and is scored on the same prompts as the reference,
|
|
4
|
+
so comparisons are paired: each case contributes one (reference, candidate) observation.
|
|
5
|
+
Pairing removes the case-to-case difficulty variation that dominates small suites, which is
|
|
6
|
+
why these tests separate models that an unpaired test calls "within noise".
|
|
7
|
+
|
|
8
|
+
Everything here is pure and deterministic. Bootstrap resampling uses its own seeded
|
|
9
|
+
generator, so the same input always yields the same interval.
|
|
10
|
+
"""
|
|
11
|
+
|
|
12
|
+
from __future__ import annotations
|
|
13
|
+
|
|
14
|
+
import math
|
|
15
|
+
import random
|
|
16
|
+
from collections.abc import Sequence
|
|
17
|
+
from dataclasses import dataclass
|
|
18
|
+
from statistics import NormalDist
|
|
19
|
+
from typing import Final
|
|
20
|
+
|
|
21
|
+
__all__ = [
|
|
22
|
+
"DEFAULT_RESAMPLES",
|
|
23
|
+
"Interval",
|
|
24
|
+
"bootstrap_mean",
|
|
25
|
+
"cases_to_bound_loss",
|
|
26
|
+
"items_to_bound_below",
|
|
27
|
+
"mcnemar_exact",
|
|
28
|
+
"paired_proportion_diff",
|
|
29
|
+
]
|
|
30
|
+
|
|
31
|
+
DEFAULT_RESAMPLES: Final = 2000
|
|
32
|
+
_CONFIDENCE: Final = 0.95
|
|
33
|
+
_Z_TWO_SIDED: Final = NormalDist().inv_cdf(1 - (1 - _CONFIDENCE) / 2)
|
|
34
|
+
_MAX_CASES_SEARCHED: Final = 1_000_000
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
@dataclass(frozen=True, slots=True)
|
|
38
|
+
class Interval:
|
|
39
|
+
"""A point estimate with a 95% confidence interval, in the metric's own units."""
|
|
40
|
+
|
|
41
|
+
estimate: float
|
|
42
|
+
low: float
|
|
43
|
+
high: float
|
|
44
|
+
|
|
45
|
+
@property
|
|
46
|
+
def excludes_zero(self) -> bool:
|
|
47
|
+
return self.low > 0.0 or self.high < 0.0
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def mcnemar_exact(ref: Sequence[bool], cand: Sequence[bool]) -> float:
|
|
51
|
+
"""Two-sided exact McNemar p-value for paired pass/fail outcomes.
|
|
52
|
+
|
|
53
|
+
Only discordant pairs carry information: b cases the reference passed and the candidate
|
|
54
|
+
failed, c the other way round. Under equal pass rates each discordant pair is a fair
|
|
55
|
+
coin, so the p-value is the two-sided binomial tail of min(b, c) in b + c trials,
|
|
56
|
+
capped at 1. With no discordant pairs there is no evidence of a difference: 1.0.
|
|
57
|
+
"""
|
|
58
|
+
_check_paired(ref, cand)
|
|
59
|
+
lost = sum(1 for r, c in zip(ref, cand, strict=True) if r and not c)
|
|
60
|
+
gained = sum(1 for r, c in zip(ref, cand, strict=True) if c and not r)
|
|
61
|
+
trials = lost + gained
|
|
62
|
+
if trials == 0:
|
|
63
|
+
return 1.0
|
|
64
|
+
tail = sum(math.comb(trials, k) for k in range(min(lost, gained) + 1))
|
|
65
|
+
# Integer division of exact big ints keeps full precision even for thousands of trials.
|
|
66
|
+
return min(1.0, 2 * tail / (1 << trials))
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
def paired_proportion_diff(ref: Sequence[bool], cand: Sequence[bool]) -> Interval:
|
|
70
|
+
"""Candidate minus reference pass rate, in percentage points, with a 95% interval.
|
|
71
|
+
|
|
72
|
+
The interval is Newcombe's hybrid score method for paired proportions (method 10 in
|
|
73
|
+
Newcombe, Statistics in Medicine 17:2635, 1998). It combines the Wilson score interval
|
|
74
|
+
of each rate with the observed correlation between the paired outcomes, stays inside
|
|
75
|
+
[-100, 100], and keeps sensible coverage at 0% and 100% where the Wald interval
|
|
76
|
+
collapses to a point.
|
|
77
|
+
"""
|
|
78
|
+
_check_paired(ref, cand)
|
|
79
|
+
if not ref:
|
|
80
|
+
raise ValueError("paired_proportion_diff needs at least one pair")
|
|
81
|
+
return _newcombe(*_table(ref, cand))
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
def bootstrap_mean(
|
|
85
|
+
values: Sequence[float],
|
|
86
|
+
*,
|
|
87
|
+
resamples: int = DEFAULT_RESAMPLES,
|
|
88
|
+
seed: int = 0,
|
|
89
|
+
) -> Interval:
|
|
90
|
+
"""Mean of `values` with a 95% percentile bootstrap interval. Deterministic for a seed."""
|
|
91
|
+
if not values:
|
|
92
|
+
raise ValueError("bootstrap_mean needs at least one value")
|
|
93
|
+
if resamples < 1:
|
|
94
|
+
raise ValueError("resamples must be at least 1")
|
|
95
|
+
n = len(values)
|
|
96
|
+
estimate = math.fsum(values) / n
|
|
97
|
+
rng = random.Random(seed) # noqa: S311 - seeded for reproducible intervals, not secrecy
|
|
98
|
+
means = sorted(
|
|
99
|
+
math.fsum(values[rng.randrange(n)] for _ in range(n)) / n for _ in range(resamples)
|
|
100
|
+
)
|
|
101
|
+
tail = (1 - _CONFIDENCE) / 2
|
|
102
|
+
return Interval(
|
|
103
|
+
estimate=estimate,
|
|
104
|
+
low=min(estimate, _percentile(means, tail)),
|
|
105
|
+
high=max(estimate, _percentile(means, 1 - tail)),
|
|
106
|
+
)
|
|
107
|
+
|
|
108
|
+
|
|
109
|
+
def items_to_bound_below(interval: Interval, items: int, bound: float) -> int | None:
|
|
110
|
+
"""Roughly how many items would put the upper end of `interval` below `bound`.
|
|
111
|
+
|
|
112
|
+
`interval` is a 95% interval of a mean over `items` items. Its upper half-width shrinks
|
|
113
|
+
with the square root of the item count, so the estimate keeps the observed mean and
|
|
114
|
+
spread and scales the half-width by sqrt(items / n). None when the observed mean is
|
|
115
|
+
already at or above `bound`, since more items would only narrow the interval around it.
|
|
116
|
+
"""
|
|
117
|
+
if items < 1:
|
|
118
|
+
raise ValueError("items must be at least 1")
|
|
119
|
+
room = bound - interval.estimate
|
|
120
|
+
if room <= 0.0:
|
|
121
|
+
return None
|
|
122
|
+
half_width = interval.high - interval.estimate
|
|
123
|
+
return max(items, math.floor(items * (half_width / room) ** 2) + 1)
|
|
124
|
+
|
|
125
|
+
|
|
126
|
+
def cases_to_bound_loss(ref: Sequence[bool], cand: Sequence[bool], margin: float) -> int | None:
|
|
127
|
+
"""Roughly how many paired cases would put the lower end of the pass-rate difference
|
|
128
|
+
(candidate minus reference, in points) at or above -`margin` points.
|
|
129
|
+
|
|
130
|
+
The observed shares of the four outcomes (both pass, only the reference passes, only
|
|
131
|
+
the candidate passes, neither) are held fixed while the case count grows, and the
|
|
132
|
+
Newcombe interval of paired_proportion_diff is recomputed until its lower bound clears
|
|
133
|
+
the margin. None when the observed difference is already at or below -`margin`, since
|
|
134
|
+
more cases would only narrow the interval around it.
|
|
135
|
+
"""
|
|
136
|
+
_check_paired(ref, cand)
|
|
137
|
+
if not ref:
|
|
138
|
+
raise ValueError("cases_to_bound_loss needs at least one pair")
|
|
139
|
+
if margin <= 0.0:
|
|
140
|
+
raise ValueError("margin must be positive")
|
|
141
|
+
n = len(ref)
|
|
142
|
+
shares = [count / n for count in _table(ref, cand)]
|
|
143
|
+
|
|
144
|
+
def clears(cases: int) -> bool:
|
|
145
|
+
both, lost, gained, neither = (share * cases for share in shares)
|
|
146
|
+
return _newcombe(both, lost, gained, neither).low >= -margin
|
|
147
|
+
|
|
148
|
+
if _newcombe(*shares).estimate <= -margin:
|
|
149
|
+
return None
|
|
150
|
+
if clears(n):
|
|
151
|
+
return n
|
|
152
|
+
low, high = n, 2 * n
|
|
153
|
+
while not clears(high):
|
|
154
|
+
if high >= _MAX_CASES_SEARCHED:
|
|
155
|
+
return None
|
|
156
|
+
low, high = high, 2 * high
|
|
157
|
+
while high - low > 1:
|
|
158
|
+
middle = (low + high) // 2
|
|
159
|
+
if clears(middle):
|
|
160
|
+
high = middle
|
|
161
|
+
else:
|
|
162
|
+
low = middle
|
|
163
|
+
return high
|
|
164
|
+
|
|
165
|
+
|
|
166
|
+
def _table(ref: Sequence[bool], cand: Sequence[bool]) -> tuple[int, int, int, int]:
|
|
167
|
+
"""Counts of (both pass, only the reference passes, only the candidate passes, neither)."""
|
|
168
|
+
both = sum(1 for r, c in zip(ref, cand, strict=True) if r and c)
|
|
169
|
+
lost = sum(1 for r, c in zip(ref, cand, strict=True) if r and not c)
|
|
170
|
+
gained = sum(1 for r, c in zip(ref, cand, strict=True) if c and not r)
|
|
171
|
+
return both, lost, gained, len(ref) - both - lost - gained
|
|
172
|
+
|
|
173
|
+
|
|
174
|
+
def _newcombe(both: float, lost: float, gained: float, neither: float) -> Interval:
|
|
175
|
+
"""Newcombe's hybrid score interval for a paired difference, from (possibly scaled)
|
|
176
|
+
cell counts of the 2x2 table, in percentage points."""
|
|
177
|
+
n = both + lost + gained + neither
|
|
178
|
+
p_ref, p_cand = (both + lost) / n, (both + gained) / n
|
|
179
|
+
low_ref, high_ref = _wilson(both + lost, n)
|
|
180
|
+
low_cand, high_cand = _wilson(both + gained, n)
|
|
181
|
+
margins = (both + lost) * (gained + neither) * (both + gained) * (lost + neither)
|
|
182
|
+
phi = 0.0 if margins == 0 else (both * neither - lost * gained) / math.sqrt(margins)
|
|
183
|
+
estimate = p_cand - p_ref
|
|
184
|
+
below = math.sqrt(
|
|
185
|
+
max(
|
|
186
|
+
0.0,
|
|
187
|
+
(p_cand - low_cand) ** 2
|
|
188
|
+
- 2 * phi * (p_cand - low_cand) * (high_ref - p_ref)
|
|
189
|
+
+ (high_ref - p_ref) ** 2,
|
|
190
|
+
)
|
|
191
|
+
)
|
|
192
|
+
above = math.sqrt(
|
|
193
|
+
max(
|
|
194
|
+
0.0,
|
|
195
|
+
(high_cand - p_cand) ** 2
|
|
196
|
+
- 2 * phi * (high_cand - p_cand) * (p_ref - low_ref)
|
|
197
|
+
+ (p_ref - low_ref) ** 2,
|
|
198
|
+
)
|
|
199
|
+
)
|
|
200
|
+
return Interval(
|
|
201
|
+
estimate=estimate * 100,
|
|
202
|
+
low=max(-1.0, estimate - below) * 100,
|
|
203
|
+
high=min(1.0, estimate + above) * 100,
|
|
204
|
+
)
|
|
205
|
+
|
|
206
|
+
|
|
207
|
+
def _wilson(successes: float, n: float) -> tuple[float, float]:
|
|
208
|
+
p = successes / n
|
|
209
|
+
z2 = _Z_TWO_SIDED * _Z_TWO_SIDED
|
|
210
|
+
centre = (p + z2 / (2 * n)) / (1 + z2 / n)
|
|
211
|
+
half = _Z_TWO_SIDED * math.sqrt(p * (1 - p) / n + z2 / (4 * n * n)) / (1 + z2 / n)
|
|
212
|
+
return max(0.0, centre - half), min(1.0, centre + half)
|
|
213
|
+
|
|
214
|
+
|
|
215
|
+
def _percentile(ordered: Sequence[float], q: float) -> float:
|
|
216
|
+
"""Linear interpolation between closest ranks, as numpy's default percentile does."""
|
|
217
|
+
position = q * (len(ordered) - 1)
|
|
218
|
+
lower = math.floor(position)
|
|
219
|
+
upper = min(lower + 1, len(ordered) - 1)
|
|
220
|
+
weight = position - lower
|
|
221
|
+
return ordered[lower] * (1 - weight) + ordered[upper] * weight
|
|
222
|
+
|
|
223
|
+
|
|
224
|
+
def _check_paired(a: Sequence[object], b: Sequence[object]) -> None:
|
|
225
|
+
if len(a) != len(b):
|
|
226
|
+
raise ValueError(f"paired samples differ in length: {len(a)} and {len(b)}")
|
|
@@ -0,0 +1,462 @@
|
|
|
1
|
+
"""Built-in prompt suites and strict loaders for user prompt files.
|
|
2
|
+
|
|
3
|
+
Suites are JSON Lines files: one task case (or scoring prompt) per line. Built-in
|
|
4
|
+
suites ship inside the package and are read with importlib.resources, so they load the
|
|
5
|
+
same way from a source checkout and from an installed wheel.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import hashlib
|
|
11
|
+
import json
|
|
12
|
+
import keyword
|
|
13
|
+
import logging
|
|
14
|
+
import os
|
|
15
|
+
import re
|
|
16
|
+
from collections.abc import Callable, Collection, Mapping, Sequence
|
|
17
|
+
from dataclasses import replace
|
|
18
|
+
from importlib import resources
|
|
19
|
+
from pathlib import Path
|
|
20
|
+
from typing import Final, Protocol, TypeVar
|
|
21
|
+
|
|
22
|
+
from quantdiff.errors import SuiteError
|
|
23
|
+
from quantdiff.metrics.jsonschema import SUPPORTED_KEYWORDS, check_schema
|
|
24
|
+
from quantdiff.types import JSONValue, Message, Role, ScoringPrompt, TaskCase, TaskKind, ToolSpec
|
|
25
|
+
|
|
26
|
+
__all__ = [
|
|
27
|
+
"BUILTIN_SUITES",
|
|
28
|
+
"MAX_CASES",
|
|
29
|
+
"MAX_FILE_BYTES",
|
|
30
|
+
"SCHEMA_KEYWORDS",
|
|
31
|
+
"case_to_dict",
|
|
32
|
+
"load_builtin",
|
|
33
|
+
"load_cases_file",
|
|
34
|
+
"load_scoring_prompts",
|
|
35
|
+
"scoring_prompts_from_cases",
|
|
36
|
+
"suite_digest",
|
|
37
|
+
]
|
|
38
|
+
|
|
39
|
+
logger = logging.getLogger(__name__)
|
|
40
|
+
|
|
41
|
+
BUILTIN_SUITES: Final[tuple[str, ...]] = ("json", "tools", "code", "chat")
|
|
42
|
+
MAX_FILE_BYTES: Final = 20 * 1024 * 1024
|
|
43
|
+
MAX_CASES: Final = 10_000
|
|
44
|
+
|
|
45
|
+
_DATA_PACKAGE: Final = "quantdiff.suites"
|
|
46
|
+
_DEFAULT_MAX_TOKENS: Final = 512
|
|
47
|
+
_MAX_MAX_TOKENS: Final = 32_768
|
|
48
|
+
_ID_PATTERN: Final = re.compile(r"[A-Za-z0-9][A-Za-z0-9._-]{0,127}")
|
|
49
|
+
_TOOL_NAME_PATTERN: Final = re.compile(r"[A-Za-z_][A-Za-z0-9_-]{0,63}")
|
|
50
|
+
_ROLES: Final[Mapping[str, Role]] = {"system": "system", "user": "user", "assistant": "assistant"}
|
|
51
|
+
_KINDS: Final[Mapping[str, TaskKind]] = {
|
|
52
|
+
"json": "json",
|
|
53
|
+
"tools": "tools",
|
|
54
|
+
"code": "code",
|
|
55
|
+
"chat": "chat",
|
|
56
|
+
}
|
|
57
|
+
_COMMON_FIELDS: Final = frozenset({"id", "kind", "messages", "max_tokens"})
|
|
58
|
+
_KIND_REQUIRED: Final[Mapping[TaskKind, frozenset[str]]] = {
|
|
59
|
+
"json": frozenset({"json_schema"}),
|
|
60
|
+
"tools": frozenset({"tools", "expected_tool"}),
|
|
61
|
+
"code": frozenset({"entry_point", "tests"}),
|
|
62
|
+
"chat": frozenset(),
|
|
63
|
+
}
|
|
64
|
+
_KIND_OPTIONAL: Final[Mapping[TaskKind, frozenset[str]]] = {
|
|
65
|
+
"json": frozenset(),
|
|
66
|
+
"tools": frozenset({"expected_arguments"}),
|
|
67
|
+
"code": frozenset(),
|
|
68
|
+
"chat": frozenset(),
|
|
69
|
+
}
|
|
70
|
+
_PROMPT_FIELDS: Final = frozenset({"prompt", "id", "system", "max_tokens"})
|
|
71
|
+
_MESSAGES_FIELDS: Final = frozenset({"messages", "id", "max_tokens"})
|
|
72
|
+
_SCORING_FIELDS: Final = frozenset({"id", "text"})
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
class _Identified(Protocol):
|
|
76
|
+
@property
|
|
77
|
+
def id(self) -> str: ...
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
_T = TypeVar("_T", bound=_Identified)
|
|
81
|
+
Record = dict[str, JSONValue]
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
# Public API ------------------------------------------------------------------------------
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
def load_builtin(name: str) -> tuple[TaskCase, ...]:
|
|
88
|
+
"""Load one of the suites listed in BUILTIN_SUITES."""
|
|
89
|
+
if name not in BUILTIN_SUITES:
|
|
90
|
+
raise SuiteError(f"unknown suite {name!r}; choose from: {', '.join(BUILTIN_SUITES)}")
|
|
91
|
+
return _parse_records(_builtin_text(f"{name}.jsonl"), f"built-in suite {name!r}", _build_case)
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
def load_cases_file(path: str | os.PathLike[str]) -> tuple[TaskCase, ...]:
|
|
95
|
+
"""Load and validate a user prompts file.
|
|
96
|
+
|
|
97
|
+
Besides full task cases, each line may be the shorthand `{"prompt": "..."}` (with
|
|
98
|
+
optional "id", "system" and "max_tokens") or `{"messages": [...]}`; both become chat
|
|
99
|
+
cases.
|
|
100
|
+
"""
|
|
101
|
+
file_path = Path(path)
|
|
102
|
+
return _parse_records(_read_text(file_path), str(file_path), _build_case)
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
def load_scoring_prompts(path: str | os.PathLike[str] | None = None) -> tuple[ScoringPrompt, ...]:
|
|
106
|
+
"""Load raw-completion prompts for logit metrics; None selects the built-in set.
|
|
107
|
+
|
|
108
|
+
Each line is `{"text": "..."}` with an optional "id", or any line a prompts file
|
|
109
|
+
accepts, in which case the last user message is scored as raw text. So one file can
|
|
110
|
+
serve as both --prompts and --scoring-prompts.
|
|
111
|
+
"""
|
|
112
|
+
if path is None:
|
|
113
|
+
return _parse_records(
|
|
114
|
+
_builtin_text("scoring.jsonl"), "built-in scoring prompts", _build_scoring_prompt
|
|
115
|
+
)
|
|
116
|
+
file_path = Path(path)
|
|
117
|
+
return _parse_records(_read_text(file_path), str(file_path), _build_scoring_prompt)
|
|
118
|
+
|
|
119
|
+
|
|
120
|
+
def scoring_prompts_from_cases(cases: Sequence[TaskCase]) -> tuple[ScoringPrompt, ...]:
|
|
121
|
+
"""Score each case's last user message as raw text, keeping the case id."""
|
|
122
|
+
return tuple(_scoring_prompt_from_case(case) for case in cases)
|
|
123
|
+
|
|
124
|
+
|
|
125
|
+
def case_to_dict(case: TaskCase) -> dict[str, JSONValue]:
|
|
126
|
+
"""Return the canonical file form of `case`; load_cases_file accepts it unchanged."""
|
|
127
|
+
out: dict[str, JSONValue] = {
|
|
128
|
+
"id": case.id,
|
|
129
|
+
"kind": case.kind,
|
|
130
|
+
"messages": [{"role": m.role, "content": m.content} for m in case.messages],
|
|
131
|
+
"max_tokens": case.max_tokens,
|
|
132
|
+
}
|
|
133
|
+
if case.json_schema is not None:
|
|
134
|
+
out["json_schema"] = case.json_schema
|
|
135
|
+
if case.tools:
|
|
136
|
+
out["tools"] = [
|
|
137
|
+
{"name": t.name, "description": t.description, "parameters": t.parameters}
|
|
138
|
+
for t in case.tools
|
|
139
|
+
]
|
|
140
|
+
if case.kind == "tools":
|
|
141
|
+
out["expected_tool"] = case.expected_tool
|
|
142
|
+
if case.expected_arguments is not None:
|
|
143
|
+
out["expected_arguments"] = case.expected_arguments
|
|
144
|
+
if case.entry_point is not None:
|
|
145
|
+
out["entry_point"] = case.entry_point
|
|
146
|
+
if case.tests is not None:
|
|
147
|
+
out["tests"] = case.tests
|
|
148
|
+
return out
|
|
149
|
+
|
|
150
|
+
|
|
151
|
+
def suite_digest(cases: Sequence[TaskCase], scoring: Sequence[ScoringPrompt]) -> str:
|
|
152
|
+
"""Return a sha256 hex digest that changes whenever any case or prompt changes."""
|
|
153
|
+
payload = {
|
|
154
|
+
"cases": [case_to_dict(case) for case in cases],
|
|
155
|
+
"scoring": [{"id": prompt.id, "text": prompt.text} for prompt in scoring],
|
|
156
|
+
}
|
|
157
|
+
canonical = json.dumps(payload, sort_keys=True, separators=(",", ":"), ensure_ascii=False)
|
|
158
|
+
return hashlib.sha256(canonical.encode("utf-8")).hexdigest()
|
|
159
|
+
|
|
160
|
+
|
|
161
|
+
# Reading ---------------------------------------------------------------------------------
|
|
162
|
+
|
|
163
|
+
|
|
164
|
+
def _builtin_text(filename: str) -> str:
|
|
165
|
+
return resources.files(_DATA_PACKAGE).joinpath("data").joinpath(filename).read_text("utf-8")
|
|
166
|
+
|
|
167
|
+
|
|
168
|
+
def _read_text(path: Path) -> str:
|
|
169
|
+
try:
|
|
170
|
+
if not path.exists():
|
|
171
|
+
raise SuiteError(f"prompts file not found: {path}")
|
|
172
|
+
if not path.is_file():
|
|
173
|
+
raise SuiteError(f"{path} is not a file")
|
|
174
|
+
if path.stat().st_size > MAX_FILE_BYTES:
|
|
175
|
+
raise SuiteError(f"{path} is larger than {MAX_FILE_BYTES // (1024 * 1024)} MB")
|
|
176
|
+
data = path.read_bytes()
|
|
177
|
+
except OSError as exc:
|
|
178
|
+
raise SuiteError(f"cannot read {path}: {exc.strerror or exc}") from None
|
|
179
|
+
try:
|
|
180
|
+
return data.decode("utf-8-sig")
|
|
181
|
+
except UnicodeDecodeError as exc:
|
|
182
|
+
raise SuiteError(f"{path} is not valid UTF-8 (byte offset {exc.start})") from None
|
|
183
|
+
|
|
184
|
+
|
|
185
|
+
def _parse_records(text: str, source: str, build: Callable[[Record, int], _T]) -> tuple[_T, ...]:
|
|
186
|
+
items: list[_T] = []
|
|
187
|
+
first_seen: dict[str, int] = {}
|
|
188
|
+
# str.splitlines would also split on U+2028 and friends, which JSON allows raw in strings.
|
|
189
|
+
for number, line in enumerate(text.split("\n"), start=1):
|
|
190
|
+
if not line.strip():
|
|
191
|
+
continue
|
|
192
|
+
if len(items) == MAX_CASES:
|
|
193
|
+
raise SuiteError(f"{source} has more than {MAX_CASES} entries")
|
|
194
|
+
try:
|
|
195
|
+
item = build(_decode_object(line), number)
|
|
196
|
+
except SuiteError as exc:
|
|
197
|
+
raise SuiteError(f"{source}, line {number}: {exc}") from None
|
|
198
|
+
if item.id in first_seen:
|
|
199
|
+
raise SuiteError(
|
|
200
|
+
f"{source}, line {number}: duplicate id {item.id!r} "
|
|
201
|
+
f"(first used on line {first_seen[item.id]})"
|
|
202
|
+
)
|
|
203
|
+
first_seen[item.id] = number
|
|
204
|
+
items.append(item)
|
|
205
|
+
if not items:
|
|
206
|
+
raise SuiteError(f"{source} contains no entries")
|
|
207
|
+
logger.debug("loaded %d entries from %s", len(items), source)
|
|
208
|
+
return tuple(items)
|
|
209
|
+
|
|
210
|
+
|
|
211
|
+
def _decode_object(line: str) -> Record:
|
|
212
|
+
try:
|
|
213
|
+
value = json.loads(line, object_pairs_hook=_unique_keys, parse_constant=_reject_constant)
|
|
214
|
+
except json.JSONDecodeError as exc:
|
|
215
|
+
raise SuiteError(f"invalid JSON: {exc.msg} (column {exc.colno})") from None
|
|
216
|
+
if not isinstance(value, dict):
|
|
217
|
+
raise SuiteError("each line must be a JSON object")
|
|
218
|
+
return value
|
|
219
|
+
|
|
220
|
+
|
|
221
|
+
def _unique_keys(pairs: list[tuple[str, JSONValue]]) -> Record:
|
|
222
|
+
record: Record = {}
|
|
223
|
+
for key, value in pairs:
|
|
224
|
+
if key in record:
|
|
225
|
+
raise SuiteError(f"duplicate key {key!r}")
|
|
226
|
+
record[key] = value
|
|
227
|
+
return record
|
|
228
|
+
|
|
229
|
+
|
|
230
|
+
def _reject_constant(name: str) -> JSONValue:
|
|
231
|
+
raise SuiteError(f"{name} is not valid JSON")
|
|
232
|
+
|
|
233
|
+
|
|
234
|
+
# Task cases ------------------------------------------------------------------------------
|
|
235
|
+
|
|
236
|
+
|
|
237
|
+
def _build_case(record: Record, line: int) -> TaskCase:
|
|
238
|
+
if "kind" in record:
|
|
239
|
+
return _full_case(record)
|
|
240
|
+
if "prompt" in record:
|
|
241
|
+
return _prompt_case(record, line)
|
|
242
|
+
if "messages" in record:
|
|
243
|
+
return _messages_case(record, line)
|
|
244
|
+
raise SuiteError('expected a "kind", "prompt" or "messages" field')
|
|
245
|
+
|
|
246
|
+
|
|
247
|
+
def _full_case(record: Record) -> TaskCase:
|
|
248
|
+
kind = _kind(record["kind"])
|
|
249
|
+
required = frozenset({"id", "kind", "messages"}) | _KIND_REQUIRED[kind]
|
|
250
|
+
_check_fields(
|
|
251
|
+
record, allowed=_COMMON_FIELDS | required | _KIND_OPTIONAL[kind], required=required
|
|
252
|
+
)
|
|
253
|
+
base = TaskCase(
|
|
254
|
+
id=_case_id(record["id"]),
|
|
255
|
+
kind=kind,
|
|
256
|
+
messages=_messages(record["messages"]),
|
|
257
|
+
max_tokens=_max_tokens(record.get("max_tokens", _DEFAULT_MAX_TOKENS)),
|
|
258
|
+
)
|
|
259
|
+
if kind == "json":
|
|
260
|
+
return _with_json_schema(base, record)
|
|
261
|
+
if kind == "tools":
|
|
262
|
+
return _with_tools(base, record)
|
|
263
|
+
if kind == "code":
|
|
264
|
+
return _with_code(base, record)
|
|
265
|
+
return base
|
|
266
|
+
|
|
267
|
+
|
|
268
|
+
def _prompt_case(record: Record, line: int) -> TaskCase:
|
|
269
|
+
_check_fields(record, allowed=_PROMPT_FIELDS, required={"prompt"})
|
|
270
|
+
messages = [Message(role="user", content=_text(record["prompt"], "prompt"))]
|
|
271
|
+
if "system" in record:
|
|
272
|
+
messages.insert(0, Message(role="system", content=_text(record["system"], "system")))
|
|
273
|
+
return TaskCase(
|
|
274
|
+
id=_case_id(record.get("id", f"prompt-{line:03d}")),
|
|
275
|
+
kind="chat",
|
|
276
|
+
messages=tuple(messages),
|
|
277
|
+
max_tokens=_max_tokens(record.get("max_tokens", _DEFAULT_MAX_TOKENS)),
|
|
278
|
+
)
|
|
279
|
+
|
|
280
|
+
|
|
281
|
+
def _messages_case(record: Record, line: int) -> TaskCase:
|
|
282
|
+
_check_fields(record, allowed=_MESSAGES_FIELDS, required={"messages"})
|
|
283
|
+
return TaskCase(
|
|
284
|
+
id=_case_id(record.get("id", f"prompt-{line:03d}")),
|
|
285
|
+
kind="chat",
|
|
286
|
+
messages=_messages(record["messages"]),
|
|
287
|
+
max_tokens=_max_tokens(record.get("max_tokens", _DEFAULT_MAX_TOKENS)),
|
|
288
|
+
)
|
|
289
|
+
|
|
290
|
+
|
|
291
|
+
def _with_json_schema(base: TaskCase, record: Record) -> TaskCase:
|
|
292
|
+
schema = _object_schema(record["json_schema"], "json_schema")
|
|
293
|
+
return replace(base, json_schema=schema)
|
|
294
|
+
|
|
295
|
+
|
|
296
|
+
def _with_tools(base: TaskCase, record: Record) -> TaskCase:
|
|
297
|
+
tools = _tools(record["tools"])
|
|
298
|
+
expected_tool = record["expected_tool"]
|
|
299
|
+
expected_arguments = record.get("expected_arguments")
|
|
300
|
+
if expected_tool is None:
|
|
301
|
+
if expected_arguments is not None:
|
|
302
|
+
raise SuiteError('"expected_arguments" requires a non-null "expected_tool"')
|
|
303
|
+
return replace(base, tools=tools)
|
|
304
|
+
by_name = {tool.name: tool for tool in tools}
|
|
305
|
+
if not isinstance(expected_tool, str) or expected_tool not in by_name:
|
|
306
|
+
raise SuiteError(f'"expected_tool" must be null or one of: {", ".join(by_name)}')
|
|
307
|
+
if expected_arguments is not None:
|
|
308
|
+
_check_expected_arguments(expected_arguments, by_name[expected_tool])
|
|
309
|
+
return replace(
|
|
310
|
+
base, tools=tools, expected_tool=expected_tool, expected_arguments=expected_arguments
|
|
311
|
+
)
|
|
312
|
+
|
|
313
|
+
|
|
314
|
+
def _with_code(base: TaskCase, record: Record) -> TaskCase:
|
|
315
|
+
entry_point = record["entry_point"]
|
|
316
|
+
if (
|
|
317
|
+
not isinstance(entry_point, str)
|
|
318
|
+
or not entry_point.isidentifier()
|
|
319
|
+
or keyword.iskeyword(entry_point)
|
|
320
|
+
):
|
|
321
|
+
raise SuiteError('"entry_point" must be a valid Python function name')
|
|
322
|
+
return replace(base, entry_point=entry_point, tests=_text(record["tests"], "tests"))
|
|
323
|
+
|
|
324
|
+
|
|
325
|
+
def _kind(value: JSONValue) -> TaskKind:
|
|
326
|
+
kind = _KINDS.get(value) if isinstance(value, str) else None
|
|
327
|
+
if kind is None:
|
|
328
|
+
raise SuiteError(f'"kind" must be one of: {", ".join(_KINDS)}')
|
|
329
|
+
return kind
|
|
330
|
+
|
|
331
|
+
|
|
332
|
+
def _case_id(value: JSONValue) -> str:
|
|
333
|
+
if not isinstance(value, str) or not _ID_PATTERN.fullmatch(value):
|
|
334
|
+
raise SuiteError(
|
|
335
|
+
'"id" must be 1 to 128 characters of letters, digits, ".", "_" or "-", '
|
|
336
|
+
"starting with a letter or digit"
|
|
337
|
+
)
|
|
338
|
+
return value
|
|
339
|
+
|
|
340
|
+
|
|
341
|
+
def _max_tokens(value: JSONValue) -> int:
|
|
342
|
+
if isinstance(value, int) and not isinstance(value, bool) and 1 <= value <= _MAX_MAX_TOKENS:
|
|
343
|
+
return value
|
|
344
|
+
raise SuiteError(f'"max_tokens" must be an integer from 1 to {_MAX_MAX_TOKENS}')
|
|
345
|
+
|
|
346
|
+
|
|
347
|
+
def _text(value: JSONValue, field: str) -> str:
|
|
348
|
+
if not isinstance(value, str) or not value.strip():
|
|
349
|
+
raise SuiteError(f'"{field}" must be a non-empty string')
|
|
350
|
+
return value
|
|
351
|
+
|
|
352
|
+
|
|
353
|
+
def _messages(value: JSONValue) -> tuple[Message, ...]:
|
|
354
|
+
if not isinstance(value, list) or not value:
|
|
355
|
+
raise SuiteError('"messages" must be a non-empty array')
|
|
356
|
+
messages = tuple(_message(item, index) for index, item in enumerate(value))
|
|
357
|
+
if messages[-1].role != "user":
|
|
358
|
+
raise SuiteError('the last message must have role "user"')
|
|
359
|
+
return messages
|
|
360
|
+
|
|
361
|
+
|
|
362
|
+
def _message(item: JSONValue, index: int) -> Message:
|
|
363
|
+
where = f"messages[{index}]"
|
|
364
|
+
if not isinstance(item, dict):
|
|
365
|
+
raise SuiteError(f"{where} must be an object")
|
|
366
|
+
_check_fields(item, allowed={"role", "content"}, required={"role", "content"}, where=where)
|
|
367
|
+
role = _ROLES.get(item["role"]) if isinstance(item["role"], str) else None
|
|
368
|
+
if role is None:
|
|
369
|
+
raise SuiteError(f"{where}.role must be one of: {', '.join(_ROLES)}")
|
|
370
|
+
return Message(role=role, content=_text(item["content"], f"{where}.content"))
|
|
371
|
+
|
|
372
|
+
|
|
373
|
+
def _check_fields(
|
|
374
|
+
record: Mapping[str, JSONValue],
|
|
375
|
+
*,
|
|
376
|
+
allowed: Collection[str],
|
|
377
|
+
required: Collection[str],
|
|
378
|
+
where: str = "",
|
|
379
|
+
) -> None:
|
|
380
|
+
prefix = f"{where}: " if where else ""
|
|
381
|
+
unknown = sorted(set(record) - set(allowed))
|
|
382
|
+
if unknown:
|
|
383
|
+
raise SuiteError(f"{prefix}unknown field {unknown[0]!r}")
|
|
384
|
+
missing = sorted(set(required) - set(record))
|
|
385
|
+
if missing:
|
|
386
|
+
raise SuiteError(f"{prefix}missing field {missing[0]!r}")
|
|
387
|
+
|
|
388
|
+
|
|
389
|
+
# Tools -----------------------------------------------------------------------------------
|
|
390
|
+
|
|
391
|
+
|
|
392
|
+
def _tools(value: JSONValue) -> tuple[ToolSpec, ...]:
|
|
393
|
+
if not isinstance(value, list) or not value:
|
|
394
|
+
raise SuiteError('"tools" must be a non-empty array')
|
|
395
|
+
tools = tuple(_tool(item, index) for index, item in enumerate(value))
|
|
396
|
+
names = [tool.name for tool in tools]
|
|
397
|
+
duplicates = sorted({name for name in names if names.count(name) > 1})
|
|
398
|
+
if duplicates:
|
|
399
|
+
raise SuiteError(f"duplicate tool name {duplicates[0]!r}")
|
|
400
|
+
return tools
|
|
401
|
+
|
|
402
|
+
|
|
403
|
+
def _tool(item: JSONValue, index: int) -> ToolSpec:
|
|
404
|
+
where = f"tools[{index}]"
|
|
405
|
+
if not isinstance(item, dict):
|
|
406
|
+
raise SuiteError(f"{where} must be an object")
|
|
407
|
+
fields = {"name", "description", "parameters"}
|
|
408
|
+
_check_fields(item, allowed=fields, required=fields, where=where)
|
|
409
|
+
name = item["name"]
|
|
410
|
+
if not isinstance(name, str) or not _TOOL_NAME_PATTERN.fullmatch(name):
|
|
411
|
+
raise SuiteError(f"{where}.name must be a function name such as get_weather")
|
|
412
|
+
if not isinstance(item["description"], str):
|
|
413
|
+
raise SuiteError(f"{where}.description must be a string")
|
|
414
|
+
parameters = _object_schema(item["parameters"], f"{where}.parameters")
|
|
415
|
+
if parameters.get("type") != "object":
|
|
416
|
+
raise SuiteError(f'{where}.parameters must have "type": "object"')
|
|
417
|
+
return ToolSpec(name=name, description=item["description"], parameters=parameters)
|
|
418
|
+
|
|
419
|
+
|
|
420
|
+
def _check_expected_arguments(value: JSONValue, tool: ToolSpec) -> None:
|
|
421
|
+
if not isinstance(value, dict):
|
|
422
|
+
raise SuiteError('"expected_arguments" must be an object or null')
|
|
423
|
+
properties = tool.parameters.get("properties", {})
|
|
424
|
+
unknown = sorted(set(value) - set(properties))
|
|
425
|
+
if unknown:
|
|
426
|
+
raise SuiteError(f"expected_arguments key {unknown[0]!r} is not a parameter of {tool.name}")
|
|
427
|
+
|
|
428
|
+
|
|
429
|
+
# JSON Schema subset ----------------------------------------------------------------------
|
|
430
|
+
|
|
431
|
+
|
|
432
|
+
SCHEMA_KEYWORDS: Final = SUPPORTED_KEYWORDS
|
|
433
|
+
"""JSON Schema keywords the task metrics can check. Anything else is rejected up front."""
|
|
434
|
+
|
|
435
|
+
|
|
436
|
+
def _object_schema(value: JSONValue, where: str) -> dict[str, JSONValue]:
|
|
437
|
+
if not isinstance(value, dict):
|
|
438
|
+
raise SuiteError(f"{where} must be a JSON object")
|
|
439
|
+
check_schema(value, where)
|
|
440
|
+
return value
|
|
441
|
+
|
|
442
|
+
|
|
443
|
+
# Scoring prompts -------------------------------------------------------------------------
|
|
444
|
+
|
|
445
|
+
|
|
446
|
+
def _build_scoring_prompt(record: Record, line: int) -> ScoringPrompt:
|
|
447
|
+
if "text" in record:
|
|
448
|
+
_check_fields(record, allowed=_SCORING_FIELDS, required={"text"})
|
|
449
|
+
return ScoringPrompt(
|
|
450
|
+
id=_case_id(record.get("id", f"score-{line:03d}")),
|
|
451
|
+
text=_text(record["text"], "text"),
|
|
452
|
+
)
|
|
453
|
+
if not record.keys() & {"kind", "prompt", "messages"}:
|
|
454
|
+
raise SuiteError('expected a "text" or "prompt" field')
|
|
455
|
+
return _scoring_prompt_from_case(_build_case(record, line))
|
|
456
|
+
|
|
457
|
+
|
|
458
|
+
def _scoring_prompt_from_case(case: TaskCase) -> ScoringPrompt:
|
|
459
|
+
user_texts = [message.content for message in case.messages if message.role == "user"]
|
|
460
|
+
if not user_texts:
|
|
461
|
+
raise SuiteError(f"case {case.id!r} has no user message to score")
|
|
462
|
+
return ScoringPrompt(id=case.id, text=user_texts[-1])
|