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
|
@@ -0,0 +1,214 @@
|
|
|
1
|
+
"""Teacher-forced logit metrics: top-1 agreement and a lower bound on KL divergence.
|
|
2
|
+
|
|
3
|
+
Servers return only the k most likely tokens at each position, so the full-vocabulary
|
|
4
|
+
KL(P_ref || Q_cand) cannot be computed. Both distributions are instead collapsed onto a
|
|
5
|
+
coarser partition: every token listed in both top-k lists keeps its own cell, the
|
|
6
|
+
reference's tokens missing from the candidate's list share a cell A, and everything else
|
|
7
|
+
shares a cell B. Merging outcomes can never increase KL divergence (data processing
|
|
8
|
+
inequality), so KL on this partition is a lower bound on the true KL.
|
|
9
|
+
|
|
10
|
+
The candidate's mass on A is not reported, but it is bounded twice over: an unlisted token
|
|
11
|
+
cannot be more likely than the least likely listed one, so Q(A) <= |A| * min listed
|
|
12
|
+
probability, and A lies outside the candidate's list, so Q(A) <= 1 - the listed mass.
|
|
13
|
+
We take the Q(A) within those bounds that minimizes the partition KL, which keeps the
|
|
14
|
+
result a lower bound while still penalizing a candidate whose top-k misses tokens the
|
|
15
|
+
reference considers likely.
|
|
16
|
+
"""
|
|
17
|
+
|
|
18
|
+
from __future__ import annotations
|
|
19
|
+
|
|
20
|
+
import math
|
|
21
|
+
import sys
|
|
22
|
+
from collections.abc import Collection, Hashable, Sequence
|
|
23
|
+
from dataclasses import dataclass
|
|
24
|
+
from typing import Final
|
|
25
|
+
|
|
26
|
+
from quantdiff.types import (
|
|
27
|
+
LogitMetrics,
|
|
28
|
+
PromptLogit,
|
|
29
|
+
ReferenceTrace,
|
|
30
|
+
TokenProb,
|
|
31
|
+
TokenStep,
|
|
32
|
+
TopK,
|
|
33
|
+
)
|
|
34
|
+
|
|
35
|
+
MASS_FLOOR: Final = 1e-12
|
|
36
|
+
"""Smallest mass a cell may hold, so a cell one side assigns ~0 cannot produce log(0)."""
|
|
37
|
+
|
|
38
|
+
PERCENTILE: Final = 0.99
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def partition_kld(ref_top: TopK, cand_top: TopK, *, by_id: bool) -> float:
|
|
42
|
+
"""Return a lower bound, in nats, on the full-vocabulary KL(P_ref || Q_cand).
|
|
43
|
+
|
|
44
|
+
Tokens are matched by id when `by_id` is True and every token carries an id,
|
|
45
|
+
otherwise by string. `cand_top` must not be empty.
|
|
46
|
+
"""
|
|
47
|
+
if not cand_top:
|
|
48
|
+
raise ValueError("candidate top-k is empty; the KL bound needs at least one token")
|
|
49
|
+
use_ids = by_id and _all_have_ids(ref_top, cand_top)
|
|
50
|
+
ref_mass = _mass_by_token(ref_top, use_ids=use_ids)
|
|
51
|
+
cand_mass = _mass_by_token(cand_top, use_ids=use_ids)
|
|
52
|
+
shared = ref_mass.keys() & cand_mass.keys()
|
|
53
|
+
ref_only = ref_mass.keys() - shared
|
|
54
|
+
|
|
55
|
+
p_shared = math.fsum(ref_mass[key] for key in shared)
|
|
56
|
+
p_missing = math.fsum(ref_mass[key] for key in ref_only)
|
|
57
|
+
p_rest = max(1.0 - p_shared - p_missing, 0.0)
|
|
58
|
+
q_unlisted = _rest_mass([cand_mass[key] for key in shared])
|
|
59
|
+
# Tokens in A are absent from the candidate's whole list, so they share what that
|
|
60
|
+
# list leaves over, not merely what the shared tokens leave over.
|
|
61
|
+
q_free = _rest_mass(cand_mass.values())
|
|
62
|
+
|
|
63
|
+
# Q(A) that minimizes the A and B terms is proportional to P; clamp it to its bounds.
|
|
64
|
+
proportional = q_unlisted * p_missing / max(p_missing + p_rest, MASS_FLOOR)
|
|
65
|
+
q_missing = min(proportional, len(ref_only) * min(cand_mass.values()), q_free, q_unlisted)
|
|
66
|
+
q_rest = q_unlisted - q_missing
|
|
67
|
+
|
|
68
|
+
cells = [(ref_mass[key], cand_mass[key]) for key in shared]
|
|
69
|
+
cells += [(p_missing, q_missing), (p_rest, q_rest)]
|
|
70
|
+
divergence = math.fsum(_kl_term(p, q) for p, q in cells)
|
|
71
|
+
# Floors and float rounding can leave a tiny negative residue on near-equal inputs.
|
|
72
|
+
return max(divergence, 0.0)
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def top1_match(ref_step: TokenStep, cand_top: TopK, *, by_id: bool) -> bool:
|
|
76
|
+
"""Return True if the candidate's most likely token is the reference's greedy choice."""
|
|
77
|
+
if not cand_top:
|
|
78
|
+
return False
|
|
79
|
+
best = max(cand_top, key=lambda token: token.logprob)
|
|
80
|
+
chosen = ref_step.chosen
|
|
81
|
+
if by_id and chosen.token_id is not None and best.token_id is not None:
|
|
82
|
+
return best.token_id == chosen.token_id
|
|
83
|
+
return best.token == chosen.token
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
def logit_metrics(
|
|
87
|
+
traces: Sequence[ReferenceTrace],
|
|
88
|
+
candidate: Sequence[Sequence[TopK]],
|
|
89
|
+
*,
|
|
90
|
+
exact_token_ids: bool,
|
|
91
|
+
) -> LogitMetrics:
|
|
92
|
+
"""Aggregate top-1 agreement and partition KL over every teacher-forced position.
|
|
93
|
+
|
|
94
|
+
`candidate[j][i]` is the candidate's top-k at position i of `traces[j]`. An empty
|
|
95
|
+
candidate list means the server stopped instead of predicting a token (it put
|
|
96
|
+
end-of-sequence first), so it counts as a top-1 miss and is left out of the KL
|
|
97
|
+
statistics, which need a distribution. Positions where the reference list is empty
|
|
98
|
+
are skipped. `per_prompt` repeats the counts for each trace so two candidates can be
|
|
99
|
+
compared prompt by prompt. Raises ValueError if the candidate shape does not match the
|
|
100
|
+
traces.
|
|
101
|
+
"""
|
|
102
|
+
_check_shapes(traces, candidate)
|
|
103
|
+
divergences: list[float] = []
|
|
104
|
+
per_prompt: list[PromptLogit] = []
|
|
105
|
+
for trace, cand_steps in zip(traces, candidate, strict=True):
|
|
106
|
+
prompt = _score_prompt(trace, cand_steps, by_id=exact_token_ids)
|
|
107
|
+
per_prompt.append(prompt.summary)
|
|
108
|
+
divergences.extend(prompt.divergences)
|
|
109
|
+
|
|
110
|
+
positions = sum(prompt.positions for prompt in per_prompt)
|
|
111
|
+
if positions == 0:
|
|
112
|
+
return LogitMetrics(
|
|
113
|
+
prompts=0,
|
|
114
|
+
positions=0,
|
|
115
|
+
top1_agreement=0.0,
|
|
116
|
+
kld_mean=0.0,
|
|
117
|
+
kld_p99=0.0,
|
|
118
|
+
kld_max=0.0,
|
|
119
|
+
exact_token_ids=exact_token_ids,
|
|
120
|
+
per_prompt=tuple(per_prompt),
|
|
121
|
+
)
|
|
122
|
+
ordered = sorted(divergences) or [0.0]
|
|
123
|
+
return LogitMetrics(
|
|
124
|
+
prompts=sum(prompt.positions > 0 for prompt in per_prompt),
|
|
125
|
+
positions=positions,
|
|
126
|
+
top1_agreement=sum(prompt.top1_matches for prompt in per_prompt) / positions,
|
|
127
|
+
kld_mean=math.fsum(ordered) / len(ordered),
|
|
128
|
+
kld_p99=_nearest_rank(ordered, PERCENTILE),
|
|
129
|
+
kld_max=ordered[-1],
|
|
130
|
+
exact_token_ids=exact_token_ids,
|
|
131
|
+
per_prompt=tuple(per_prompt),
|
|
132
|
+
)
|
|
133
|
+
|
|
134
|
+
|
|
135
|
+
@dataclass(frozen=True, slots=True)
|
|
136
|
+
class _ScoredPrompt:
|
|
137
|
+
summary: PromptLogit
|
|
138
|
+
divergences: list[float]
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
def _score_prompt(
|
|
142
|
+
trace: ReferenceTrace, cand_steps: Sequence[TopK], *, by_id: bool
|
|
143
|
+
) -> _ScoredPrompt:
|
|
144
|
+
divergences: list[float] = []
|
|
145
|
+
matches = 0
|
|
146
|
+
positions = 0
|
|
147
|
+
for ref_step, cand_top in zip(trace.steps, cand_steps, strict=True):
|
|
148
|
+
if not ref_step.top:
|
|
149
|
+
continue
|
|
150
|
+
positions += 1
|
|
151
|
+
matches += top1_match(ref_step, cand_top, by_id=by_id)
|
|
152
|
+
if cand_top:
|
|
153
|
+
divergences.append(partition_kld(ref_step.top, cand_top, by_id=by_id))
|
|
154
|
+
kld_mean = math.fsum(divergences) / len(divergences) if divergences else None
|
|
155
|
+
summary = PromptLogit(trace.prompt_id, positions, matches, kld_mean)
|
|
156
|
+
return _ScoredPrompt(summary, divergences)
|
|
157
|
+
|
|
158
|
+
|
|
159
|
+
def _check_shapes(traces: Sequence[ReferenceTrace], candidate: Sequence[Sequence[TopK]]) -> None:
|
|
160
|
+
if len(traces) != len(candidate):
|
|
161
|
+
raise ValueError(
|
|
162
|
+
f"got candidate scores for {len(candidate)} traces, expected {len(traces)}"
|
|
163
|
+
)
|
|
164
|
+
for trace, cand_steps in zip(traces, candidate, strict=True):
|
|
165
|
+
if len(trace.steps) != len(cand_steps):
|
|
166
|
+
raise ValueError(
|
|
167
|
+
f"trace {trace.prompt_id!r} has {len(trace.steps)} steps but the candidate "
|
|
168
|
+
f"scored {len(cand_steps)} positions"
|
|
169
|
+
)
|
|
170
|
+
|
|
171
|
+
|
|
172
|
+
def _all_have_ids(*tops: TopK) -> bool:
|
|
173
|
+
return all(token.token_id is not None for top in tops for token in top)
|
|
174
|
+
|
|
175
|
+
|
|
176
|
+
def _mass_by_token(top: TopK, *, use_ids: bool) -> dict[Hashable, float]:
|
|
177
|
+
"""Map each token key to its probability, summing duplicates of the same key."""
|
|
178
|
+
masses: dict[Hashable, float] = {}
|
|
179
|
+
for token in top:
|
|
180
|
+
key = _token_key(token, use_ids=use_ids)
|
|
181
|
+
masses[key] = masses.get(key, 0.0) + _probability(token.logprob)
|
|
182
|
+
return masses
|
|
183
|
+
|
|
184
|
+
|
|
185
|
+
def _token_key(token: TokenProb, *, use_ids: bool) -> Hashable:
|
|
186
|
+
return token.token_id if use_ids else token.token
|
|
187
|
+
|
|
188
|
+
|
|
189
|
+
def _probability(logprob: float) -> float:
|
|
190
|
+
if math.isnan(logprob):
|
|
191
|
+
raise ValueError("logprob is NaN")
|
|
192
|
+
# Servers occasionally report logprobs a hair above 0 for near-certain tokens.
|
|
193
|
+
return math.exp(min(logprob, 0.0))
|
|
194
|
+
|
|
195
|
+
|
|
196
|
+
def _rest_mass(listed: Collection[float]) -> float:
|
|
197
|
+
"""Return the candidate mass outside `listed`, padded by the rounding error of the sum.
|
|
198
|
+
|
|
199
|
+
When a cap binds at a tiny true mass, even an error of a few ulps in 1 - sum is a large
|
|
200
|
+
relative error. Padding errs toward more candidate mass, which can only lower the KL.
|
|
201
|
+
"""
|
|
202
|
+
rounding = (len(listed) + 1) * sys.float_info.epsilon
|
|
203
|
+
return min(max(1.0 - math.fsum(listed) + rounding, MASS_FLOOR), 1.0)
|
|
204
|
+
|
|
205
|
+
|
|
206
|
+
def _kl_term(p: float, q: float) -> float:
|
|
207
|
+
if p <= 0.0:
|
|
208
|
+
return 0.0
|
|
209
|
+
return p * math.log(p / max(q, MASS_FLOOR))
|
|
210
|
+
|
|
211
|
+
|
|
212
|
+
def _nearest_rank(ordered: Sequence[float], fraction: float) -> float:
|
|
213
|
+
rank = max(math.ceil(fraction * len(ordered)), 1)
|
|
214
|
+
return ordered[rank - 1]
|
|
@@ -0,0 +1,114 @@
|
|
|
1
|
+
"""Per-case scoring for task suites and the summaries built from it."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import math
|
|
6
|
+
from collections.abc import Sequence
|
|
7
|
+
from typing import Final
|
|
8
|
+
|
|
9
|
+
from quantdiff.errors import SuiteError
|
|
10
|
+
from quantdiff.metrics.codeexec import DEFAULT_TIMEOUT_SECONDS, run_code_case
|
|
11
|
+
from quantdiff.metrics.jsonschema import extract_json, validate
|
|
12
|
+
from quantdiff.metrics.toolcheck import check_tool_case
|
|
13
|
+
from quantdiff.types import CaseOutcome, ChatResult, PerfMetrics, TaskCase, TaskKind, TaskMetrics
|
|
14
|
+
|
|
15
|
+
MAX_LISTED_FAILURES: Final = 25
|
|
16
|
+
SCORED_KINDS: Final[tuple[TaskKind, ...]] = ("json", "tools", "code")
|
|
17
|
+
CODE_EXEC_DISABLED: Final = "code execution disabled (pass --allow-code-exec)"
|
|
18
|
+
CHAT_SCORED_BY_AGREEMENT: Final = "scored by agreement"
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def evaluate_case(
|
|
22
|
+
case: TaskCase,
|
|
23
|
+
result: ChatResult,
|
|
24
|
+
*,
|
|
25
|
+
allow_code_exec: bool,
|
|
26
|
+
timeout_seconds: float = DEFAULT_TIMEOUT_SECONDS,
|
|
27
|
+
) -> CaseOutcome:
|
|
28
|
+
"""Score one chat result against its case. Code cases are skipped unless allowed."""
|
|
29
|
+
if case.kind == "json":
|
|
30
|
+
return _check_json_case(case, result)
|
|
31
|
+
if case.kind == "tools":
|
|
32
|
+
return check_tool_case(case, result)
|
|
33
|
+
if case.kind == "code":
|
|
34
|
+
if not allow_code_exec:
|
|
35
|
+
return CaseOutcome(case.id, case.kind, passed=None, reason=CODE_EXEC_DISABLED)
|
|
36
|
+
return run_code_case(case, result.text, timeout_seconds=timeout_seconds)
|
|
37
|
+
return CaseOutcome(case.id, case.kind, passed=None, reason=CHAT_SCORED_BY_AGREEMENT)
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def summarize_tasks(outcomes: Sequence[CaseOutcome]) -> tuple[TaskMetrics, ...]:
|
|
41
|
+
"""Return one TaskMetrics per scored kind present, in json, tools, code order."""
|
|
42
|
+
summaries = []
|
|
43
|
+
for kind in SCORED_KINDS:
|
|
44
|
+
group = [outcome for outcome in outcomes if outcome.kind == kind]
|
|
45
|
+
if not group:
|
|
46
|
+
continue
|
|
47
|
+
failures = [outcome for outcome in group if outcome.passed is False]
|
|
48
|
+
summaries.append(
|
|
49
|
+
TaskMetrics(
|
|
50
|
+
kind=kind,
|
|
51
|
+
total=len(group),
|
|
52
|
+
passed=sum(outcome.passed is True for outcome in group),
|
|
53
|
+
skipped=sum(outcome.passed is None for outcome in group),
|
|
54
|
+
failures=tuple(failures[:MAX_LISTED_FAILURES]),
|
|
55
|
+
)
|
|
56
|
+
)
|
|
57
|
+
return tuple(summaries)
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def perf_metrics(results: Sequence[ChatResult]) -> PerfMetrics:
|
|
61
|
+
"""Aggregate generation speed and mean request latency over chat results.
|
|
62
|
+
|
|
63
|
+
When the server timed its own decoding for every result that generated more than one
|
|
64
|
+
token, `tokens_per_second` is total completion tokens over total decode time (the
|
|
65
|
+
token-weighted harmonic mean of the per-request rates) with source "server". That
|
|
66
|
+
excludes model load and prompt processing, so it is comparable across models. A
|
|
67
|
+
one-token answer has no decode step to time and is left out of that figure.
|
|
68
|
+
|
|
69
|
+
Otherwise it falls back to completion tokens over wall-clock request time with source
|
|
70
|
+
"wall_clock", which includes network overhead and prompt evaluation and is therefore
|
|
71
|
+
only comparable between candidates measured on the same machine with the same cases.
|
|
72
|
+
"""
|
|
73
|
+
mean_latency = sum(result.seconds for result in results) / len(results) if results else None
|
|
74
|
+
server_rate = _server_decode_rate(results)
|
|
75
|
+
if server_rate is not None:
|
|
76
|
+
return PerfMetrics(server_rate, mean_latency, source="server")
|
|
77
|
+
timed = [
|
|
78
|
+
result for result in results if result.completion_tokens is not None and result.seconds > 0
|
|
79
|
+
]
|
|
80
|
+
total_seconds = sum(result.seconds for result in timed)
|
|
81
|
+
tokens_per_second = (
|
|
82
|
+
sum(result.completion_tokens or 0 for result in timed) / total_seconds
|
|
83
|
+
if total_seconds > 0
|
|
84
|
+
else None
|
|
85
|
+
)
|
|
86
|
+
return PerfMetrics(tokens_per_second, mean_latency, source="wall_clock")
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
def _server_decode_rate(results: Sequence[ChatResult]) -> float | None:
|
|
90
|
+
tokens = 0
|
|
91
|
+
decode_seconds: list[float] = []
|
|
92
|
+
for result in results:
|
|
93
|
+
count = result.completion_tokens
|
|
94
|
+
if count is None or count < 2:
|
|
95
|
+
continue
|
|
96
|
+
rate = result.decode_tokens_per_second
|
|
97
|
+
if not rate:
|
|
98
|
+
return None
|
|
99
|
+
tokens += count
|
|
100
|
+
decode_seconds.append(count / rate)
|
|
101
|
+
return tokens / math.fsum(decode_seconds) if decode_seconds else None
|
|
102
|
+
|
|
103
|
+
|
|
104
|
+
def _check_json_case(case: TaskCase, result: ChatResult) -> CaseOutcome:
|
|
105
|
+
if case.json_schema is None:
|
|
106
|
+
raise SuiteError(f"json case {case.id!r} has no json_schema")
|
|
107
|
+
value, problem = extract_json(result.text)
|
|
108
|
+
if problem is not None:
|
|
109
|
+
return CaseOutcome(case.id, case.kind, passed=False, reason=problem)
|
|
110
|
+
errors = validate(value, case.json_schema)
|
|
111
|
+
if errors:
|
|
112
|
+
extra = f" (and {len(errors) - 1} more)" if len(errors) > 1 else ""
|
|
113
|
+
return CaseOutcome(case.id, case.kind, passed=False, reason=errors[0] + extra)
|
|
114
|
+
return CaseOutcome(case.id, case.kind, passed=True)
|
|
@@ -0,0 +1,66 @@
|
|
|
1
|
+
"""Text normalization and similarity for comparing free-form chat answers."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import difflib
|
|
6
|
+
import math
|
|
7
|
+
import re
|
|
8
|
+
from collections.abc import Sequence
|
|
9
|
+
from typing import Final
|
|
10
|
+
|
|
11
|
+
from quantdiff.types import AgreementMetrics
|
|
12
|
+
|
|
13
|
+
MAX_COMPARED_CHARS: Final = 4000
|
|
14
|
+
"""SequenceMatcher is quadratic in the worst case, so long answers are truncated first."""
|
|
15
|
+
|
|
16
|
+
_OPEN_TAG: Final = "<think>"
|
|
17
|
+
_CLOSE_TAG: Final = "</think>"
|
|
18
|
+
_LEADING_THINK: Final = re.compile(r"\A\s*<think>.*?</think>", re.DOTALL)
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def strip_reasoning(text: str) -> str:
|
|
22
|
+
"""Remove a leading <think>...</think> block emitted by reasoning models.
|
|
23
|
+
|
|
24
|
+
Some chat templates open the block inside the prompt, so the answer starts with the
|
|
25
|
+
reasoning and only contains the closing tag; that prefix is removed too. An opening
|
|
26
|
+
tag that is never closed means generation stopped mid-reasoning, so nothing is left.
|
|
27
|
+
"""
|
|
28
|
+
match = _LEADING_THINK.match(text)
|
|
29
|
+
if match:
|
|
30
|
+
return text[match.end() :]
|
|
31
|
+
if text.lstrip().startswith(_OPEN_TAG):
|
|
32
|
+
return ""
|
|
33
|
+
if _CLOSE_TAG in text and _OPEN_TAG not in text:
|
|
34
|
+
return text.split(_CLOSE_TAG, 1)[1]
|
|
35
|
+
return text
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def normalize(text: str) -> str:
|
|
39
|
+
"""Strip reasoning, trim, and collapse every run of whitespace to one space."""
|
|
40
|
+
return " ".join(strip_reasoning(text).split())
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
def similarity(a: str, b: str) -> float:
|
|
44
|
+
"""Return a similarity ratio in [0, 1] between two answers after normalization."""
|
|
45
|
+
left = normalize(a)[:MAX_COMPARED_CHARS]
|
|
46
|
+
right = normalize(b)[:MAX_COMPARED_CHARS]
|
|
47
|
+
if left == right:
|
|
48
|
+
return 1.0
|
|
49
|
+
return difflib.SequenceMatcher(None, left, right, autojunk=False).ratio()
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def agreement_metrics(answers: Sequence[tuple[str, str, str]]) -> AgreementMetrics:
|
|
53
|
+
"""Score (case id, reference answer, candidate answer) triples by exact match and mean
|
|
54
|
+
similarity, keeping each case's similarity for paired comparisons."""
|
|
55
|
+
if not answers:
|
|
56
|
+
return AgreementMetrics(cases=0, exact_match_rate=0.0, mean_similarity=0.0)
|
|
57
|
+
exact = sum(normalize(reference) == normalize(candidate) for _, reference, candidate in answers)
|
|
58
|
+
per_case = tuple(
|
|
59
|
+
(case_id, similarity(reference, candidate)) for case_id, reference, candidate in answers
|
|
60
|
+
)
|
|
61
|
+
return AgreementMetrics(
|
|
62
|
+
cases=len(answers),
|
|
63
|
+
exact_match_rate=exact / len(answers),
|
|
64
|
+
mean_similarity=math.fsum(score for _, score in per_case) / len(answers),
|
|
65
|
+
per_case=per_case,
|
|
66
|
+
)
|
|
@@ -0,0 +1,99 @@
|
|
|
1
|
+
"""Scoring for tool-calling cases: right tool, valid arguments, expected values."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from typing import Final
|
|
6
|
+
|
|
7
|
+
from quantdiff.errors import SuiteError
|
|
8
|
+
from quantdiff.metrics.jsonschema import validate
|
|
9
|
+
from quantdiff.types import CaseOutcome, ChatResult, JSONValue, TaskCase, ToolCall
|
|
10
|
+
|
|
11
|
+
_PREVIEW_CHARS: Final = 60
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def check_tool_case(case: TaskCase, result: ChatResult) -> CaseOutcome:
|
|
15
|
+
"""Pass when the first tool call matches the case's expectations.
|
|
16
|
+
|
|
17
|
+
A case with no `expected_tool` passes only when the model answers without calling
|
|
18
|
+
any tool. Raises SuiteError if the expected tool is not among the case's tools.
|
|
19
|
+
"""
|
|
20
|
+
reason = _failure_reason(case, result)
|
|
21
|
+
return CaseOutcome(case_id=case.id, kind=case.kind, passed=reason is None, reason=reason or "")
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def arguments_match(expected: JSONValue, actual: JSONValue) -> bool:
|
|
25
|
+
"""Lenient equality for tool arguments.
|
|
26
|
+
|
|
27
|
+
Numbers compare numerically (1 equals 1.0, but never equals true), strings compare
|
|
28
|
+
after strip() and casefold(), lists compare element by element, and objects match
|
|
29
|
+
when every expected key is present with a matching value.
|
|
30
|
+
"""
|
|
31
|
+
if isinstance(expected, bool) or isinstance(actual, bool):
|
|
32
|
+
return type(expected) is type(actual) and expected == actual
|
|
33
|
+
if isinstance(expected, (int, float)):
|
|
34
|
+
return isinstance(actual, (int, float)) and expected == actual
|
|
35
|
+
if isinstance(expected, str):
|
|
36
|
+
return isinstance(actual, str) and _fold(expected) == _fold(actual)
|
|
37
|
+
if isinstance(expected, list):
|
|
38
|
+
return (
|
|
39
|
+
isinstance(actual, list)
|
|
40
|
+
and len(expected) == len(actual)
|
|
41
|
+
and all(map(arguments_match, expected, actual))
|
|
42
|
+
)
|
|
43
|
+
if isinstance(expected, dict):
|
|
44
|
+
return isinstance(actual, dict) and all(
|
|
45
|
+
key in actual and arguments_match(value, actual[key]) for key, value in expected.items()
|
|
46
|
+
)
|
|
47
|
+
return expected is None and actual is None
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def _failure_reason(case: TaskCase, result: ChatResult) -> str | None:
|
|
51
|
+
expected_tool = case.expected_tool
|
|
52
|
+
if expected_tool is None:
|
|
53
|
+
return _direct_answer_failure(result)
|
|
54
|
+
if not result.tool_calls:
|
|
55
|
+
return f"made no tool call, expected {expected_tool}"
|
|
56
|
+
|
|
57
|
+
call = result.tool_calls[0]
|
|
58
|
+
if call.name != expected_tool:
|
|
59
|
+
return f"called {call.name}, expected {expected_tool}"
|
|
60
|
+
if call.arguments is None:
|
|
61
|
+
return f"{call.name} arguments are not a JSON object: {_preview(call.raw_arguments)}"
|
|
62
|
+
schema_errors = validate(call.arguments, _parameters_schema(case, expected_tool))
|
|
63
|
+
if schema_errors:
|
|
64
|
+
extra = f" (and {len(schema_errors) - 1} more)" if len(schema_errors) > 1 else ""
|
|
65
|
+
return f"{call.name} arguments violate the schema: {schema_errors[0]}{extra}"
|
|
66
|
+
return _argument_mismatch(call, case.expected_arguments or {})
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
def _direct_answer_failure(result: ChatResult) -> str | None:
|
|
70
|
+
if result.tool_calls:
|
|
71
|
+
return f"called {result.tool_calls[0].name}, expected a direct answer"
|
|
72
|
+
return None
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def _parameters_schema(case: TaskCase, tool_name: str) -> dict[str, JSONValue]:
|
|
76
|
+
for tool in case.tools:
|
|
77
|
+
if tool.name == tool_name:
|
|
78
|
+
return tool.parameters
|
|
79
|
+
raise SuiteError(f"case {case.id!r} expects tool {tool_name!r} but does not define it")
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
def _argument_mismatch(call: ToolCall, expected: dict[str, JSONValue]) -> str | None:
|
|
83
|
+
arguments = call.arguments or {}
|
|
84
|
+
for key, value in expected.items():
|
|
85
|
+
if key not in arguments:
|
|
86
|
+
return f"{call.name} is missing argument {key!r}"
|
|
87
|
+
if not arguments_match(value, arguments[key]):
|
|
88
|
+
actual = _preview(repr(arguments[key]))
|
|
89
|
+
return f"{call.name} argument {key!r} is {actual}, expected {_preview(repr(value))}"
|
|
90
|
+
return None
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
def _fold(text: str) -> str:
|
|
94
|
+
return text.strip().casefold()
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
def _preview(text: str) -> str:
|
|
98
|
+
flat = " ".join(text.split())
|
|
99
|
+
return flat if len(flat) <= _PREVIEW_CHARS else flat[: _PREVIEW_CHARS - 3] + "..."
|