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/runner.py
ADDED
|
@@ -0,0 +1,492 @@
|
|
|
1
|
+
"""Orchestrates a comparison run: connect to everything, run the reference, then each
|
|
2
|
+
candidate in turn.
|
|
3
|
+
|
|
4
|
+
Every server is contacted before any real work starts, so a typo in a model tag fails in
|
|
5
|
+
a second instead of after minutes of reference generation. Candidates then run one at a
|
|
6
|
+
time on purpose: they usually share a GPU, and running them in parallel would distort
|
|
7
|
+
throughput numbers and the servers' own caching. A server error on one request is
|
|
8
|
+
recorded and never discards results that were already measured.
|
|
9
|
+
"""
|
|
10
|
+
|
|
11
|
+
from __future__ import annotations
|
|
12
|
+
|
|
13
|
+
import logging
|
|
14
|
+
from collections.abc import Callable, Sequence
|
|
15
|
+
from dataclasses import dataclass, field, replace
|
|
16
|
+
from datetime import datetime, timezone
|
|
17
|
+
|
|
18
|
+
from quantdiff._version import __version__
|
|
19
|
+
from quantdiff.backends import open_backend
|
|
20
|
+
from quantdiff.backends._common import forced_texts
|
|
21
|
+
from quantdiff.backends.base import Backend
|
|
22
|
+
from quantdiff.cache import ReferenceCache, ReferenceOutputs, cache_key
|
|
23
|
+
from quantdiff.errors import BackendError, CapabilityError
|
|
24
|
+
from quantdiff.metrics import (
|
|
25
|
+
agreement_metrics,
|
|
26
|
+
evaluate_case,
|
|
27
|
+
logit_metrics,
|
|
28
|
+
perf_metrics,
|
|
29
|
+
summarize_tasks,
|
|
30
|
+
)
|
|
31
|
+
from quantdiff.metrics.tasks import SCORED_KINDS
|
|
32
|
+
from quantdiff.preflight import run_preflight
|
|
33
|
+
from quantdiff.suites import suite_digest
|
|
34
|
+
from quantdiff.types import (
|
|
35
|
+
CandidateResult,
|
|
36
|
+
CandidateSpec,
|
|
37
|
+
CaseOutcome,
|
|
38
|
+
ChatResult,
|
|
39
|
+
LogitMetrics,
|
|
40
|
+
Message,
|
|
41
|
+
PreflightFinding,
|
|
42
|
+
ProgressEvent,
|
|
43
|
+
ProgressPhase,
|
|
44
|
+
ReferenceTrace,
|
|
45
|
+
Report,
|
|
46
|
+
RunSettings,
|
|
47
|
+
ScoringPrompt,
|
|
48
|
+
ServerInfo,
|
|
49
|
+
TaskCase,
|
|
50
|
+
TokenStep,
|
|
51
|
+
TopK,
|
|
52
|
+
)
|
|
53
|
+
|
|
54
|
+
logger = logging.getLogger(__name__)
|
|
55
|
+
|
|
56
|
+
ProgressFn = Callable[[ProgressEvent], None]
|
|
57
|
+
BackendFactory = Callable[[CandidateSpec], Backend]
|
|
58
|
+
|
|
59
|
+
# Progress units are weighted by measured cost so the ETA is honest. One unit is one
|
|
60
|
+
# teacher-forced position (a single-token request on a cached prefix). On a consumer GPU a
|
|
61
|
+
# chat case costs about 30 of those, a reference scoring prompt about one per generated
|
|
62
|
+
# token, and the pre-flight long-prompt probe about as much as six cases.
|
|
63
|
+
_CASE_UNITS = 30
|
|
64
|
+
_PREFLIGHT_UNITS = 180
|
|
65
|
+
_WARMUP_UNITS = 10
|
|
66
|
+
_WARMUP_MESSAGES = (Message(role="user", content="Hi"),)
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
@dataclass(frozen=True, slots=True)
|
|
70
|
+
class RunPlan:
|
|
71
|
+
"""Everything a run needs, fully resolved and validated."""
|
|
72
|
+
|
|
73
|
+
reference: CandidateSpec
|
|
74
|
+
candidates: tuple[CandidateSpec, ...]
|
|
75
|
+
cases: tuple[TaskCase, ...]
|
|
76
|
+
scoring: tuple[ScoringPrompt, ...]
|
|
77
|
+
settings: RunSettings
|
|
78
|
+
title: str
|
|
79
|
+
hf_repo: str | None = None
|
|
80
|
+
offline: bool = False
|
|
81
|
+
preflight: bool = True
|
|
82
|
+
context_probe_tokens: int = 6000
|
|
83
|
+
code_timeout_seconds: float = 10.0
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
@dataclass(slots=True)
|
|
87
|
+
class _Progress:
|
|
88
|
+
"""Counts weighted work units and forwards ProgressEvents."""
|
|
89
|
+
|
|
90
|
+
sink: ProgressFn
|
|
91
|
+
total: int
|
|
92
|
+
completed: int = 0
|
|
93
|
+
model: str = ""
|
|
94
|
+
|
|
95
|
+
def step(self, phase: ProgressPhase, units: int = 1, detail: str = "") -> None:
|
|
96
|
+
self.completed = min(self.completed + units, self.total)
|
|
97
|
+
self.sink(ProgressEvent(phase, self.model, self.completed, self.total, detail))
|
|
98
|
+
|
|
99
|
+
def finish(self) -> None:
|
|
100
|
+
self.completed = self.total
|
|
101
|
+
self.model = ""
|
|
102
|
+
self.sink(ProgressEvent("done", "", self.total, self.total))
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
@dataclass(slots=True)
|
|
106
|
+
class _Answers:
|
|
107
|
+
"""Chat answers by case id, plus the cases the server failed to answer."""
|
|
108
|
+
|
|
109
|
+
results: dict[str, ChatResult] = field(default_factory=dict)
|
|
110
|
+
failures: list[str] = field(default_factory=list)
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
def execute(
|
|
114
|
+
plan: RunPlan,
|
|
115
|
+
*,
|
|
116
|
+
backend_factory: BackendFactory = open_backend,
|
|
117
|
+
cache: ReferenceCache | None = None,
|
|
118
|
+
progress: ProgressFn | None = None,
|
|
119
|
+
) -> Report:
|
|
120
|
+
"""Run the plan and return a Report. Pass cache=None to always rerun the reference.
|
|
121
|
+
|
|
122
|
+
Raises BackendError before doing any work if a server cannot be reached or does not
|
|
123
|
+
serve the requested model.
|
|
124
|
+
"""
|
|
125
|
+
tracker = _Progress(sink=progress or _log_progress, total=_total_units(plan))
|
|
126
|
+
backends = _connect(plan, backend_factory, tracker)
|
|
127
|
+
try:
|
|
128
|
+
reference, *candidates = backends
|
|
129
|
+
ref_backend, ref_info = reference
|
|
130
|
+
tracker.model = plan.reference.label
|
|
131
|
+
outputs, cached = _reference_outputs(plan, ref_backend, ref_info, cache, tracker)
|
|
132
|
+
ref_result = _reference_result(plan, ref_backend, ref_info, outputs, tracker, cached=cached)
|
|
133
|
+
results = tuple(
|
|
134
|
+
_run_candidate(
|
|
135
|
+
plan,
|
|
136
|
+
spec,
|
|
137
|
+
backend=backend,
|
|
138
|
+
info=info,
|
|
139
|
+
reference=ref_backend,
|
|
140
|
+
ref_info=ref_info,
|
|
141
|
+
outputs=outputs,
|
|
142
|
+
tracker=tracker,
|
|
143
|
+
)
|
|
144
|
+
for spec, (backend, info) in zip(plan.candidates, candidates, strict=True)
|
|
145
|
+
)
|
|
146
|
+
finally:
|
|
147
|
+
for backend, _ in backends:
|
|
148
|
+
backend.close()
|
|
149
|
+
tracker.finish()
|
|
150
|
+
|
|
151
|
+
return Report(
|
|
152
|
+
quantdiff_version=__version__,
|
|
153
|
+
created_at=datetime.now(timezone.utc).isoformat(timespec="seconds"),
|
|
154
|
+
title=plan.title,
|
|
155
|
+
settings=plan.settings,
|
|
156
|
+
reference=ref_result,
|
|
157
|
+
candidates=results,
|
|
158
|
+
)
|
|
159
|
+
|
|
160
|
+
|
|
161
|
+
def _total_units(plan: RunPlan) -> int:
|
|
162
|
+
preflight = _PREFLIGHT_UNITS if plan.preflight else 0
|
|
163
|
+
models = 1 + len(plan.candidates)
|
|
164
|
+
cases = _WARMUP_UNITS + _CASE_UNITS * len(plan.cases)
|
|
165
|
+
scoring = len(plan.scoring) * _scoring_units(plan)
|
|
166
|
+
return models + (cases + scoring + preflight) * models
|
|
167
|
+
|
|
168
|
+
|
|
169
|
+
def _scoring_units(plan: RunPlan) -> int:
|
|
170
|
+
"""Units for one scoring prompt: one per generated or teacher-forced token."""
|
|
171
|
+
return max(plan.settings.score_tokens, 1)
|
|
172
|
+
|
|
173
|
+
|
|
174
|
+
def _connect(
|
|
175
|
+
plan: RunPlan, backend_factory: BackendFactory, tracker: _Progress
|
|
176
|
+
) -> list[tuple[Backend, ServerInfo]]:
|
|
177
|
+
"""Open every backend and fetch its info, or raise one error naming every failure."""
|
|
178
|
+
connected: list[tuple[Backend, ServerInfo]] = []
|
|
179
|
+
problems: list[str] = []
|
|
180
|
+
for spec in (plan.reference, *plan.candidates):
|
|
181
|
+
tracker.model = spec.label
|
|
182
|
+
try:
|
|
183
|
+
backend = backend_factory(spec)
|
|
184
|
+
except BackendError as exc:
|
|
185
|
+
problems.append(f"{spec.label}: {exc}")
|
|
186
|
+
continue
|
|
187
|
+
try:
|
|
188
|
+
connected.append((backend, backend.info()))
|
|
189
|
+
except BackendError as exc:
|
|
190
|
+
backend.close()
|
|
191
|
+
problems.append(f"{spec.label}: {exc}")
|
|
192
|
+
tracker.step("connect")
|
|
193
|
+
if problems:
|
|
194
|
+
for backend, _ in connected:
|
|
195
|
+
backend.close()
|
|
196
|
+
raise BackendError("cannot start the comparison:\n " + "\n ".join(problems))
|
|
197
|
+
return connected
|
|
198
|
+
|
|
199
|
+
|
|
200
|
+
# Reference --------------------------------------------------------------------------------
|
|
201
|
+
|
|
202
|
+
|
|
203
|
+
def _reference_outputs(
|
|
204
|
+
plan: RunPlan,
|
|
205
|
+
reference: Backend,
|
|
206
|
+
info: ServerInfo,
|
|
207
|
+
cache: ReferenceCache | None,
|
|
208
|
+
tracker: _Progress,
|
|
209
|
+
) -> tuple[ReferenceOutputs, bool]:
|
|
210
|
+
"""Return the reference outputs and whether they came from the cache."""
|
|
211
|
+
key = cache_key(
|
|
212
|
+
info,
|
|
213
|
+
base_url=plan.reference.base_url,
|
|
214
|
+
suite_digest=suite_digest(plan.cases, plan.scoring),
|
|
215
|
+
top_k=plan.settings.top_k,
|
|
216
|
+
score_tokens=plan.settings.score_tokens,
|
|
217
|
+
seed=plan.settings.seed,
|
|
218
|
+
)
|
|
219
|
+
if cache is not None:
|
|
220
|
+
cached = cache.load(key)
|
|
221
|
+
if cached is not None:
|
|
222
|
+
units = (
|
|
223
|
+
_WARMUP_UNITS
|
|
224
|
+
+ _CASE_UNITS * len(plan.cases)
|
|
225
|
+
+ _scoring_units(plan) * len(plan.scoring)
|
|
226
|
+
)
|
|
227
|
+
tracker.step("reference", units, "loaded from cache")
|
|
228
|
+
return cached, True
|
|
229
|
+
|
|
230
|
+
answers = _answer_cases(plan, reference, tracker, "reference")
|
|
231
|
+
if answers.failures:
|
|
232
|
+
logger.warning("reference failed %d case(s); they are left out", len(answers.failures))
|
|
233
|
+
traces = _reference_traces(plan, reference, info, tracker)
|
|
234
|
+
outputs = ReferenceOutputs(answers=answers.results, traces=traces)
|
|
235
|
+
if cache is not None and not answers.failures:
|
|
236
|
+
cache.store(key, outputs)
|
|
237
|
+
return outputs, False
|
|
238
|
+
|
|
239
|
+
|
|
240
|
+
def _reference_traces(
|
|
241
|
+
plan: RunPlan, reference: Backend, info: ServerInfo, tracker: _Progress
|
|
242
|
+
) -> tuple[ReferenceTrace, ...]:
|
|
243
|
+
if not plan.scoring or not info.supports_logprobs:
|
|
244
|
+
tracker.step("scoring", _scoring_units(plan) * len(plan.scoring), "skipped: no logprobs")
|
|
245
|
+
return ()
|
|
246
|
+
traces = []
|
|
247
|
+
for index, prompt in enumerate(plan.scoring, start=1):
|
|
248
|
+
try:
|
|
249
|
+
steps = reference.generate_scored(
|
|
250
|
+
prompt.text, max_tokens=plan.settings.score_tokens, top_k=plan.settings.top_k
|
|
251
|
+
)
|
|
252
|
+
ids = reference.tokenize(prompt.text) if info.exact_token_ids else None
|
|
253
|
+
except CapabilityError as exc:
|
|
254
|
+
logger.warning("reference exposes no logprobs, skipping the logit tier: %s", exc)
|
|
255
|
+
remaining = len(plan.scoring) - index + 1
|
|
256
|
+
tracker.step("scoring", _scoring_units(plan) * remaining, "skipped: no logprobs")
|
|
257
|
+
return ()
|
|
258
|
+
except BackendError as exc:
|
|
259
|
+
logger.warning("reference failed scoring prompt %s: %s", prompt.id, exc)
|
|
260
|
+
else:
|
|
261
|
+
traces.append(ReferenceTrace(prompt.id, ids, tuple(steps)))
|
|
262
|
+
tracker.step("scoring", _scoring_units(plan), f"{index}/{len(plan.scoring)}")
|
|
263
|
+
return tuple(traces)
|
|
264
|
+
|
|
265
|
+
|
|
266
|
+
def _reference_result(
|
|
267
|
+
plan: RunPlan,
|
|
268
|
+
reference: Backend,
|
|
269
|
+
info: ServerInfo,
|
|
270
|
+
outputs: ReferenceOutputs,
|
|
271
|
+
tracker: _Progress,
|
|
272
|
+
*,
|
|
273
|
+
cached: bool,
|
|
274
|
+
) -> CandidateResult:
|
|
275
|
+
findings = _preflight(plan, reference, None, tracker)
|
|
276
|
+
outcomes = _score_cases(plan, outputs.answers)
|
|
277
|
+
perf = perf_metrics(list(outputs.answers.values()))
|
|
278
|
+
return CandidateResult(
|
|
279
|
+
spec=plan.reference,
|
|
280
|
+
info=info,
|
|
281
|
+
logit=None,
|
|
282
|
+
tasks=summarize_tasks(outcomes),
|
|
283
|
+
agreement=None,
|
|
284
|
+
# Timings from an earlier run stay visible but labelled, since the machine's load
|
|
285
|
+
# may differ from the fresh candidate runs.
|
|
286
|
+
perf=replace(perf, source="cached") if cached else perf,
|
|
287
|
+
preflight=findings,
|
|
288
|
+
outcomes=outcomes,
|
|
289
|
+
)
|
|
290
|
+
|
|
291
|
+
|
|
292
|
+
# Candidates -------------------------------------------------------------------------------
|
|
293
|
+
|
|
294
|
+
|
|
295
|
+
def _run_candidate(
|
|
296
|
+
plan: RunPlan,
|
|
297
|
+
spec: CandidateSpec,
|
|
298
|
+
*,
|
|
299
|
+
backend: Backend,
|
|
300
|
+
info: ServerInfo,
|
|
301
|
+
reference: Backend,
|
|
302
|
+
ref_info: ServerInfo,
|
|
303
|
+
outputs: ReferenceOutputs,
|
|
304
|
+
tracker: _Progress,
|
|
305
|
+
) -> CandidateResult:
|
|
306
|
+
tracker.model = spec.label
|
|
307
|
+
findings = _preflight(plan, backend, reference, tracker)
|
|
308
|
+
answers = _answer_cases(plan, backend, tracker, spec.label)
|
|
309
|
+
logit, logit_error = _logit_tier(
|
|
310
|
+
plan,
|
|
311
|
+
backend,
|
|
312
|
+
info=info,
|
|
313
|
+
ref_info=ref_info,
|
|
314
|
+
outputs=outputs,
|
|
315
|
+
findings=findings,
|
|
316
|
+
tracker=tracker,
|
|
317
|
+
)
|
|
318
|
+
|
|
319
|
+
errors = []
|
|
320
|
+
if answers.failures:
|
|
321
|
+
errors.append(_case_failure_summary(answers.failures))
|
|
322
|
+
if logit_error is not None:
|
|
323
|
+
errors.append(logit_error)
|
|
324
|
+
chats = [
|
|
325
|
+
(case.id, outputs.answers[case.id].text, answers.results[case.id].text)
|
|
326
|
+
for case in plan.cases
|
|
327
|
+
if case.kind == "chat" and case.id in answers.results and case.id in outputs.answers
|
|
328
|
+
]
|
|
329
|
+
outcomes = _score_cases(plan, answers.results)
|
|
330
|
+
return CandidateResult(
|
|
331
|
+
spec=spec,
|
|
332
|
+
info=info,
|
|
333
|
+
logit=logit,
|
|
334
|
+
tasks=summarize_tasks(outcomes),
|
|
335
|
+
agreement=agreement_metrics(chats) if chats else None,
|
|
336
|
+
perf=perf_metrics(list(answers.results.values())),
|
|
337
|
+
preflight=findings,
|
|
338
|
+
errors=tuple(errors),
|
|
339
|
+
outcomes=outcomes,
|
|
340
|
+
)
|
|
341
|
+
|
|
342
|
+
|
|
343
|
+
def _logit_tier(
|
|
344
|
+
plan: RunPlan,
|
|
345
|
+
backend: Backend,
|
|
346
|
+
*,
|
|
347
|
+
info: ServerInfo,
|
|
348
|
+
ref_info: ServerInfo,
|
|
349
|
+
outputs: ReferenceOutputs,
|
|
350
|
+
findings: Sequence[PreflightFinding],
|
|
351
|
+
tracker: _Progress,
|
|
352
|
+
) -> tuple[LogitMetrics | None, str | None]:
|
|
353
|
+
"""Return logit metrics, or None plus a reason when the tier cannot run fairly."""
|
|
354
|
+
units = len(plan.scoring) * _scoring_units(plan)
|
|
355
|
+
if not outputs.traces or not info.supports_logprobs:
|
|
356
|
+
tracker.step("scoring", units, "skipped")
|
|
357
|
+
return None, None
|
|
358
|
+
if any(f.check == "tokenizer" and f.severity == "fail" for f in findings):
|
|
359
|
+
tracker.step("scoring", units, "skipped")
|
|
360
|
+
return None, "logit tier skipped: tokenizer differs from the reference"
|
|
361
|
+
return _teacher_force(
|
|
362
|
+
plan, backend, info=info, ref_info=ref_info, outputs=outputs, tracker=tracker
|
|
363
|
+
)
|
|
364
|
+
|
|
365
|
+
|
|
366
|
+
def _teacher_force(
|
|
367
|
+
plan: RunPlan,
|
|
368
|
+
backend: Backend,
|
|
369
|
+
*,
|
|
370
|
+
info: ServerInfo,
|
|
371
|
+
ref_info: ServerInfo,
|
|
372
|
+
outputs: ReferenceOutputs,
|
|
373
|
+
tracker: _Progress,
|
|
374
|
+
) -> tuple[LogitMetrics | None, str | None]:
|
|
375
|
+
exact = ref_info.exact_token_ids and info.exact_token_ids
|
|
376
|
+
prompts = {prompt.id: prompt.text for prompt in plan.scoring}
|
|
377
|
+
traces = (
|
|
378
|
+
outputs.traces
|
|
379
|
+
if exact
|
|
380
|
+
else tuple(_text_scorable(trace, prompts[trace.prompt_id]) for trace in outputs.traces)
|
|
381
|
+
)
|
|
382
|
+
budget = len(plan.scoring) * _scoring_units(plan)
|
|
383
|
+
spent = 0
|
|
384
|
+
scored: list[list[TopK]] = []
|
|
385
|
+
for index, trace in enumerate(traces, start=1):
|
|
386
|
+
try:
|
|
387
|
+
tops = backend.score_continuation(
|
|
388
|
+
prompts[trace.prompt_id],
|
|
389
|
+
trace.steps,
|
|
390
|
+
top_k=plan.settings.top_k,
|
|
391
|
+
prompt_token_ids=trace.prompt_token_ids if exact else None,
|
|
392
|
+
)
|
|
393
|
+
except BackendError as exc:
|
|
394
|
+
tracker.step("scoring", budget - spent, "failed")
|
|
395
|
+
return None, f"logit tier failed: {exc}"
|
|
396
|
+
scored.append(tops)
|
|
397
|
+
spent += _scoring_units(plan)
|
|
398
|
+
tracker.step("scoring", _scoring_units(plan), f"{index}/{len(traces)}")
|
|
399
|
+
tracker.step("scoring", budget - spent)
|
|
400
|
+
return logit_metrics(traces, scored, exact_token_ids=exact), None
|
|
401
|
+
|
|
402
|
+
|
|
403
|
+
def _text_scorable(trace: ReferenceTrace, prompt: str) -> ReferenceTrace:
|
|
404
|
+
"""The trace with every position that text cannot reproduce marked unscored.
|
|
405
|
+
|
|
406
|
+
A reference scored by token ids can stop inside a multi-byte character, and no text
|
|
407
|
+
prompt ends there. Marking those positions unscored (an empty `top`) makes metrics skip
|
|
408
|
+
them for this candidate instead of counting the empty result as a top-1 miss.
|
|
409
|
+
"""
|
|
410
|
+
texts = forced_texts(prompt, trace.steps)
|
|
411
|
+
steps = tuple(
|
|
412
|
+
step if text is not None else TokenStep(step.chosen, ())
|
|
413
|
+
for step, text in zip(trace.steps, texts, strict=True)
|
|
414
|
+
)
|
|
415
|
+
return ReferenceTrace(trace.prompt_id, trace.prompt_token_ids, steps)
|
|
416
|
+
|
|
417
|
+
|
|
418
|
+
# Shared steps -----------------------------------------------------------------------------
|
|
419
|
+
|
|
420
|
+
|
|
421
|
+
def _answer_cases(plan: RunPlan, backend: Backend, tracker: _Progress, label: str) -> _Answers:
|
|
422
|
+
_warm_up(plan, backend, tracker, label)
|
|
423
|
+
answers = _Answers()
|
|
424
|
+
total = len(plan.cases)
|
|
425
|
+
for index, case in enumerate(plan.cases, start=1):
|
|
426
|
+
try:
|
|
427
|
+
answers.results[case.id] = backend.chat(
|
|
428
|
+
case.messages,
|
|
429
|
+
max_tokens=case.max_tokens,
|
|
430
|
+
tools=case.tools,
|
|
431
|
+
json_schema=case.json_schema if case.kind == "json" else None,
|
|
432
|
+
seed=plan.settings.seed,
|
|
433
|
+
)
|
|
434
|
+
except BackendError as exc:
|
|
435
|
+
logger.warning("%s failed case %s: %s", label, case.id, exc)
|
|
436
|
+
answers.failures.append(f"{case.id}: {exc}")
|
|
437
|
+
tracker.step("cases", _CASE_UNITS, f"{index}/{total}")
|
|
438
|
+
return answers
|
|
439
|
+
|
|
440
|
+
|
|
441
|
+
def _warm_up(plan: RunPlan, backend: Backend, tracker: _Progress, label: str) -> None:
|
|
442
|
+
"""Send one tiny request so model load time never lands in the first case's timing.
|
|
443
|
+
|
|
444
|
+
A failure here is only logged: if the server is really broken, the cases report it.
|
|
445
|
+
"""
|
|
446
|
+
try:
|
|
447
|
+
backend.chat(_WARMUP_MESSAGES, max_tokens=1, seed=plan.settings.seed)
|
|
448
|
+
except BackendError as exc:
|
|
449
|
+
logger.debug("%s warm-up request failed: %s", label, exc)
|
|
450
|
+
tracker.step("cases", _WARMUP_UNITS, "warm-up")
|
|
451
|
+
|
|
452
|
+
|
|
453
|
+
def _case_failure_summary(failures: Sequence[str]) -> str:
|
|
454
|
+
first = failures[0]
|
|
455
|
+
more = f" (and {len(failures) - 1} more)" if len(failures) > 1 else ""
|
|
456
|
+
return f"{len(failures)} case(s) failed with server errors and were not scored: {first}{more}"
|
|
457
|
+
|
|
458
|
+
|
|
459
|
+
def _score_cases(plan: RunPlan, answers: dict[str, ChatResult]) -> tuple[CaseOutcome, ...]:
|
|
460
|
+
"""Outcomes of the answered json, tools and code cases, in plan order."""
|
|
461
|
+
return tuple(
|
|
462
|
+
evaluate_case(
|
|
463
|
+
case,
|
|
464
|
+
answers[case.id],
|
|
465
|
+
allow_code_exec=plan.settings.allow_code_exec,
|
|
466
|
+
timeout_seconds=plan.code_timeout_seconds,
|
|
467
|
+
)
|
|
468
|
+
for case in plan.cases
|
|
469
|
+
if case.id in answers and case.kind in SCORED_KINDS
|
|
470
|
+
)
|
|
471
|
+
|
|
472
|
+
|
|
473
|
+
def _preflight(
|
|
474
|
+
plan: RunPlan, backend: Backend, reference: Backend | None, tracker: _Progress
|
|
475
|
+
) -> tuple[PreflightFinding, ...]:
|
|
476
|
+
if not plan.preflight:
|
|
477
|
+
return ()
|
|
478
|
+
findings = run_preflight(
|
|
479
|
+
backend,
|
|
480
|
+
reference=reference,
|
|
481
|
+
hf_repo=plan.hf_repo,
|
|
482
|
+
offline=plan.offline,
|
|
483
|
+
context_probe_tokens=plan.context_probe_tokens,
|
|
484
|
+
)
|
|
485
|
+
tracker.step("preflight", _PREFLIGHT_UNITS)
|
|
486
|
+
return findings
|
|
487
|
+
|
|
488
|
+
|
|
489
|
+
def _log_progress(event: ProgressEvent) -> None:
|
|
490
|
+
logger.info(
|
|
491
|
+
"%s %s %d/%d %s", event.phase, event.model, event.completed, event.total, event.detail
|
|
492
|
+
)
|
quantdiff/spec.py
ADDED
|
@@ -0,0 +1,154 @@
|
|
|
1
|
+
"""Parse candidate spec strings from the command line into CandidateSpec values.
|
|
2
|
+
|
|
3
|
+
Grammar::
|
|
4
|
+
|
|
5
|
+
spec = [label "="] target
|
|
6
|
+
target = "ollama:" tag
|
|
7
|
+
| "llamacpp:" base_url
|
|
8
|
+
| "openai:" base_url "#" model ["@env:" env_var]
|
|
9
|
+
|
|
10
|
+
label = any non-empty text before the first "=" that contains no ":"
|
|
11
|
+
env_var = name of an environment variable holding an API key ([A-Za-z_][A-Za-z0-9_]*)
|
|
12
|
+
|
|
13
|
+
Examples::
|
|
14
|
+
|
|
15
|
+
ollama:qwen2.5:0.5b-instruct-q8_0
|
|
16
|
+
q4=ollama:qwen2.5:7b-instruct-q4_K_M
|
|
17
|
+
llamacpp:http://127.0.0.1:8080
|
|
18
|
+
openai:http://127.0.0.1:1234/v1#qwen2.5-7b-instruct
|
|
19
|
+
vllm=openai:http://gpu-box:8000/v1#Qwen/Qwen2.5-7B-Instruct@env:VLLM_API_KEY
|
|
20
|
+
|
|
21
|
+
Ollama specs use OLLAMA_HOST (`host`, `host:port` or a full URL) when it is set, and
|
|
22
|
+
http://127.0.0.1:11434 otherwise. Default labels are the Ollama tag, `llamacpp@host:port`
|
|
23
|
+
and the OpenAI model name. Making labels unique across candidates is the caller's job.
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
from __future__ import annotations
|
|
27
|
+
|
|
28
|
+
import dataclasses
|
|
29
|
+
import os
|
|
30
|
+
import re
|
|
31
|
+
import urllib.parse
|
|
32
|
+
from typing import Final
|
|
33
|
+
|
|
34
|
+
from quantdiff._http import validate_base_url
|
|
35
|
+
from quantdiff.errors import BackendError, SpecError
|
|
36
|
+
from quantdiff.types import CandidateSpec
|
|
37
|
+
|
|
38
|
+
DEFAULT_OLLAMA_URL: Final = "http://127.0.0.1:11434"
|
|
39
|
+
_OLLAMA_DEFAULT_PORT: Final = 11434
|
|
40
|
+
_API_KEY_MARKER: Final = "@env:"
|
|
41
|
+
_ENV_NAME: Final = re.compile(r"[A-Za-z_][A-Za-z0-9_]*")
|
|
42
|
+
_OLLAMA_TAG: Final = re.compile(r"[A-Za-z0-9][A-Za-z0-9._:/-]*")
|
|
43
|
+
_BIND_ALL_HOSTS: Final = frozenset({"0.0.0.0", "::"}) # noqa: S104 - matched, not bound
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
def parse_spec(text: str) -> CandidateSpec:
|
|
47
|
+
"""Parse one candidate spec string. Raises SpecError if it is malformed."""
|
|
48
|
+
stripped = text.strip()
|
|
49
|
+
if not stripped:
|
|
50
|
+
raise SpecError("empty candidate spec")
|
|
51
|
+
label, target = _split_label(stripped)
|
|
52
|
+
kind, separator, rest = target.partition(":")
|
|
53
|
+
if not separator or not rest:
|
|
54
|
+
raise SpecError(f"{text!r}: expected <kind>:<target>, for example ollama:<tag>")
|
|
55
|
+
if kind == "ollama":
|
|
56
|
+
spec = _ollama_spec(rest)
|
|
57
|
+
elif kind == "llamacpp":
|
|
58
|
+
spec = _llamacpp_spec(rest)
|
|
59
|
+
elif kind == "openai":
|
|
60
|
+
spec = _openai_spec(rest)
|
|
61
|
+
else:
|
|
62
|
+
raise SpecError(f"{text!r}: unknown backend {kind!r}; use ollama, llamacpp or openai")
|
|
63
|
+
return spec if label is None else dataclasses.replace(spec, label=label)
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
def ollama_base_url(host: str | None) -> str:
|
|
67
|
+
"""Resolve an OLLAMA_HOST value the way the Ollama CLI does."""
|
|
68
|
+
value = (host or "").strip()
|
|
69
|
+
if not value:
|
|
70
|
+
return DEFAULT_OLLAMA_URL
|
|
71
|
+
url = value if "://" in value else f"http://{value}"
|
|
72
|
+
parsed = urllib.parse.urlsplit(url)
|
|
73
|
+
hostname = parsed.hostname
|
|
74
|
+
if not hostname:
|
|
75
|
+
raise SpecError(f"OLLAMA_HOST={value!r} has no host")
|
|
76
|
+
# OLLAMA_HOST is often a listen address such as 0.0.0.0, which clients cannot dial.
|
|
77
|
+
if hostname in _BIND_ALL_HOSTS:
|
|
78
|
+
hostname = "127.0.0.1"
|
|
79
|
+
port = _port(parsed, value)
|
|
80
|
+
if port is None and "://" not in value:
|
|
81
|
+
port = _OLLAMA_DEFAULT_PORT
|
|
82
|
+
netloc = f"[{hostname}]" if ":" in hostname else hostname
|
|
83
|
+
if port is not None:
|
|
84
|
+
netloc = f"{netloc}:{port}"
|
|
85
|
+
return _base_url(urllib.parse.urlunsplit((parsed.scheme, netloc, parsed.path, "", "")))
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
def _split_label(text: str) -> tuple[str | None, str]:
|
|
89
|
+
head, separator, tail = text.partition("=")
|
|
90
|
+
if not separator or ":" in head:
|
|
91
|
+
return None, text
|
|
92
|
+
label = head.strip()
|
|
93
|
+
if not label:
|
|
94
|
+
raise SpecError(f"{text!r}: the label before '=' is empty")
|
|
95
|
+
if not label.isprintable():
|
|
96
|
+
raise SpecError(f"{text!r}: the label contains control characters")
|
|
97
|
+
return label, tail.strip()
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
def _ollama_spec(tag: str) -> CandidateSpec:
|
|
101
|
+
if not _OLLAMA_TAG.fullmatch(tag):
|
|
102
|
+
raise SpecError(f"invalid Ollama tag {tag!r}")
|
|
103
|
+
return CandidateSpec(
|
|
104
|
+
kind="ollama",
|
|
105
|
+
base_url=ollama_base_url(os.environ.get("OLLAMA_HOST")),
|
|
106
|
+
model=tag,
|
|
107
|
+
label=tag,
|
|
108
|
+
)
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
def _llamacpp_spec(url: str) -> CandidateSpec:
|
|
112
|
+
base_url = _base_url(url)
|
|
113
|
+
return CandidateSpec(
|
|
114
|
+
kind="llamacpp",
|
|
115
|
+
base_url=base_url,
|
|
116
|
+
model="",
|
|
117
|
+
label=f"llamacpp@{urllib.parse.urlsplit(base_url).netloc}",
|
|
118
|
+
)
|
|
119
|
+
|
|
120
|
+
|
|
121
|
+
def _openai_spec(rest: str) -> CandidateSpec:
|
|
122
|
+
url, separator, model_part = rest.partition("#")
|
|
123
|
+
if not separator:
|
|
124
|
+
raise SpecError(f"openai spec {rest!r} needs a model: openai:<base_url>#<model>")
|
|
125
|
+
model, marker, env_name = model_part.rpartition(_API_KEY_MARKER)
|
|
126
|
+
if not marker:
|
|
127
|
+
model, env_name = model_part, ""
|
|
128
|
+
model = model.strip()
|
|
129
|
+
if not model or any(char.isspace() for char in model):
|
|
130
|
+
raise SpecError(f"invalid model name {model!r} in openai spec")
|
|
131
|
+
if marker and not _ENV_NAME.fullmatch(env_name):
|
|
132
|
+
raise SpecError(f"invalid environment variable name {env_name!r} after {_API_KEY_MARKER}")
|
|
133
|
+
return CandidateSpec(
|
|
134
|
+
kind="openai",
|
|
135
|
+
base_url=_base_url(url),
|
|
136
|
+
model=model,
|
|
137
|
+
label=model,
|
|
138
|
+
api_key_env=env_name or None,
|
|
139
|
+
)
|
|
140
|
+
|
|
141
|
+
|
|
142
|
+
def _base_url(url: str) -> str:
|
|
143
|
+
try:
|
|
144
|
+
return validate_base_url(url.strip())
|
|
145
|
+
except BackendError as exc:
|
|
146
|
+
# The message from validate_base_url never echoes credentials; the raw URL might.
|
|
147
|
+
raise SpecError(f"invalid server URL: {exc}") from None
|
|
148
|
+
|
|
149
|
+
|
|
150
|
+
def _port(parsed: urllib.parse.SplitResult, raw: str) -> int | None:
|
|
151
|
+
try:
|
|
152
|
+
return parsed.port
|
|
153
|
+
except ValueError:
|
|
154
|
+
raise SpecError(f"OLLAMA_HOST={raw!r} has an invalid port") from None
|