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.
Files changed (45) hide show
  1. quantdiff/__init__.py +53 -0
  2. quantdiff/__main__.py +5 -0
  3. quantdiff/_http.py +151 -0
  4. quantdiff/_text.py +13 -0
  5. quantdiff/_version.py +1 -0
  6. quantdiff/api.py +340 -0
  7. quantdiff/backends/__init__.py +28 -0
  8. quantdiff/backends/_common.py +342 -0
  9. quantdiff/backends/base.py +91 -0
  10. quantdiff/backends/llamacpp.py +428 -0
  11. quantdiff/backends/ollama.py +359 -0
  12. quantdiff/backends/openai_compat.py +338 -0
  13. quantdiff/cache.py +240 -0
  14. quantdiff/card.py +1664 -0
  15. quantdiff/cli.py +377 -0
  16. quantdiff/discover.py +488 -0
  17. quantdiff/errors.py +45 -0
  18. quantdiff/metrics/__init__.py +36 -0
  19. quantdiff/metrics/codeexec.py +428 -0
  20. quantdiff/metrics/jsonschema.py +610 -0
  21. quantdiff/metrics/logit.py +214 -0
  22. quantdiff/metrics/tasks.py +114 -0
  23. quantdiff/metrics/textsim.py +66 -0
  24. quantdiff/metrics/toolcheck.py +99 -0
  25. quantdiff/png.py +360 -0
  26. quantdiff/preflight.py +365 -0
  27. quantdiff/progress.py +283 -0
  28. quantdiff/py.typed +0 -0
  29. quantdiff/report.py +780 -0
  30. quantdiff/runner.py +492 -0
  31. quantdiff/spec.py +154 -0
  32. quantdiff/stats.py +226 -0
  33. quantdiff/suites/__init__.py +462 -0
  34. quantdiff/suites/data/chat.jsonl +22 -0
  35. quantdiff/suites/data/code.jsonl +32 -0
  36. quantdiff/suites/data/json.jsonl +34 -0
  37. quantdiff/suites/data/scoring.jsonl +41 -0
  38. quantdiff/suites/data/tools.jsonl +32 -0
  39. quantdiff/types.py +322 -0
  40. quantdiff/verdict.py +1513 -0
  41. quantdiff-0.1.0rc1.dist-info/METADATA +514 -0
  42. quantdiff-0.1.0rc1.dist-info/RECORD +45 -0
  43. quantdiff-0.1.0rc1.dist-info/WHEEL +4 -0
  44. quantdiff-0.1.0rc1.dist-info/entry_points.txt +2 -0
  45. quantdiff-0.1.0rc1.dist-info/licenses/LICENSE +202 -0
quantdiff/cache.py ADDED
@@ -0,0 +1,240 @@
1
+ """On-disk cache of reference model outputs.
2
+
3
+ Running the reference is the most expensive part of a comparison, and its outputs do not
4
+ change between runs with the same model, suite and settings. Entries are plain JSON, keyed
5
+ by a SHA-256 digest of everything that can change the output, and written atomically.
6
+ A corrupt or mismatched entry is treated as a miss, never as an error.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import base64
12
+ import hashlib
13
+ import json
14
+ import logging
15
+ import os
16
+ import sys
17
+ import tempfile
18
+ from dataclasses import dataclass
19
+ from pathlib import Path
20
+ from typing import Final
21
+
22
+ from quantdiff.types import (
23
+ ChatResult,
24
+ JSONValue,
25
+ ReferenceTrace,
26
+ ServerInfo,
27
+ TokenProb,
28
+ TokenStep,
29
+ ToolCall,
30
+ )
31
+
32
+ logger = logging.getLogger(__name__)
33
+
34
+ CACHE_FORMAT: Final = 4
35
+ """Bumped whenever the entry layout changes, so older entries are misses."""
36
+ _MAX_ENTRY_BYTES: Final = 256 * 1024 * 1024
37
+
38
+
39
+ def default_cache_dir() -> Path:
40
+ """Return the per-user cache directory, honoring QUANTDIFF_CACHE_DIR first."""
41
+ override = os.environ.get("QUANTDIFF_CACHE_DIR")
42
+ if override:
43
+ return Path(override)
44
+ if sys.platform == "win32":
45
+ base = os.environ.get("LOCALAPPDATA")
46
+ if base:
47
+ return Path(base) / "quantdiff" / "cache"
48
+ xdg = os.environ.get("XDG_CACHE_HOME")
49
+ return (Path(xdg) if xdg else Path.home() / ".cache") / "quantdiff"
50
+
51
+
52
+ @dataclass(frozen=True, slots=True)
53
+ class ReferenceOutputs:
54
+ """Everything the runner needs from the reference model."""
55
+
56
+ answers: dict[str, ChatResult]
57
+ """Chat answers keyed by TaskCase id."""
58
+ traces: tuple[ReferenceTrace, ...]
59
+ """Greedy scored continuations, empty when the reference exposes no logprobs."""
60
+
61
+
62
+ def cache_key(
63
+ info: ServerInfo,
64
+ *,
65
+ base_url: str,
66
+ suite_digest: str,
67
+ top_k: int,
68
+ score_tokens: int,
69
+ seed: int,
70
+ ) -> str:
71
+ """Digest of every input that can change the reference outputs.
72
+
73
+ `info.details` carries each backend's fingerprint of the served weights (an Ollama
74
+ digest, a GGUF file size) and the chat template is included directly, so re-pulling a
75
+ fixed upload or a patched template invalidates the entry.
76
+ """
77
+ material = {
78
+ "format": CACHE_FORMAT,
79
+ "backend": info.backend,
80
+ "base_url": base_url,
81
+ "model": info.model,
82
+ "chat_template": info.chat_template,
83
+ "details": sorted(info.details),
84
+ "exact_token_ids": info.exact_token_ids,
85
+ "suite": suite_digest,
86
+ "top_k": top_k,
87
+ "score_tokens": score_tokens,
88
+ "seed": seed,
89
+ }
90
+ encoded = json.dumps(material, sort_keys=True, separators=(",", ":")).encode("utf-8")
91
+ return hashlib.sha256(encoded).hexdigest()
92
+
93
+
94
+ class ReferenceCache:
95
+ """A directory of `<key>.json` entries."""
96
+
97
+ def __init__(self, directory: Path | None = None) -> None:
98
+ self.directory = directory if directory is not None else default_cache_dir()
99
+
100
+ def load(self, key: str) -> ReferenceOutputs | None:
101
+ path = self._path(key)
102
+ try:
103
+ if path.stat().st_size > _MAX_ENTRY_BYTES:
104
+ logger.warning("ignoring oversized cache entry %s", path)
105
+ return None
106
+ data = json.loads(path.read_text(encoding="utf-8"))
107
+ return _outputs_from_dict(data)
108
+ except FileNotFoundError:
109
+ return None
110
+ except (OSError, ValueError, KeyError, TypeError) as exc:
111
+ logger.warning("ignoring unreadable cache entry %s: %s", path, exc)
112
+ return None
113
+
114
+ def store(self, key: str, outputs: ReferenceOutputs) -> None:
115
+ self.directory.mkdir(parents=True, exist_ok=True)
116
+ payload = json.dumps(_outputs_to_dict(outputs), ensure_ascii=False, separators=(",", ":"))
117
+ fd, tmp_name = tempfile.mkstemp(dir=self.directory, prefix=".tmp-", suffix=".json")
118
+ tmp_path = Path(tmp_name)
119
+ try:
120
+ with os.fdopen(fd, "w", encoding="utf-8") as handle:
121
+ handle.write(payload)
122
+ tmp_path.replace(self._path(key))
123
+ except BaseException:
124
+ tmp_path.unlink(missing_ok=True)
125
+ raise
126
+
127
+ def _path(self, key: str) -> Path:
128
+ if len(key) != 64 or not all(char in "0123456789abcdef" for char in key):
129
+ raise ValueError(f"invalid cache key: {key!r}")
130
+ return self.directory / f"{key}.json"
131
+
132
+
133
+ # Serialization ----------------------------------------------------------------------------
134
+
135
+
136
+ def _outputs_to_dict(outputs: ReferenceOutputs) -> dict[str, JSONValue]:
137
+ return {
138
+ "format": CACHE_FORMAT,
139
+ "answers": {case_id: _chat_to_dict(result) for case_id, result in outputs.answers.items()},
140
+ "traces": [_trace_to_dict(trace) for trace in outputs.traces],
141
+ }
142
+
143
+
144
+ def _outputs_from_dict(data: JSONValue) -> ReferenceOutputs:
145
+ if not isinstance(data, dict) or data.get("format") != CACHE_FORMAT:
146
+ raise ValueError("unsupported cache format")
147
+ answers = {str(case_id): _chat_from_dict(item) for case_id, item in data["answers"].items()}
148
+ traces = tuple(_trace_from_dict(item) for item in data["traces"])
149
+ return ReferenceOutputs(answers=answers, traces=traces)
150
+
151
+
152
+ def _chat_to_dict(result: ChatResult) -> dict[str, JSONValue]:
153
+ return {
154
+ "text": result.text,
155
+ "tool_calls": [
156
+ {"name": call.name, "arguments": call.arguments, "raw_arguments": call.raw_arguments}
157
+ for call in result.tool_calls
158
+ ],
159
+ "finish_reason": result.finish_reason,
160
+ "prompt_tokens": result.prompt_tokens,
161
+ "completion_tokens": result.completion_tokens,
162
+ "seconds": result.seconds,
163
+ "decode_tokens_per_second": result.decode_tokens_per_second,
164
+ }
165
+
166
+
167
+ def _chat_from_dict(data: dict[str, JSONValue]) -> ChatResult:
168
+ calls = tuple(
169
+ ToolCall(
170
+ name=str(call["name"]),
171
+ arguments=call["arguments"] if isinstance(call["arguments"], dict) else None,
172
+ raw_arguments=str(call["raw_arguments"]),
173
+ )
174
+ for call in data["tool_calls"]
175
+ )
176
+ return ChatResult(
177
+ text=str(data["text"]),
178
+ tool_calls=calls,
179
+ finish_reason=_optional_str(data["finish_reason"]),
180
+ prompt_tokens=_optional_int(data["prompt_tokens"]),
181
+ completion_tokens=_optional_int(data["completion_tokens"]),
182
+ seconds=float(data["seconds"]),
183
+ decode_tokens_per_second=_optional_float(data["decode_tokens_per_second"]),
184
+ )
185
+
186
+
187
+ def _trace_to_dict(trace: ReferenceTrace) -> dict[str, JSONValue]:
188
+ return {
189
+ "prompt_id": trace.prompt_id,
190
+ "prompt_token_ids": None
191
+ if trace.prompt_token_ids is None
192
+ else list(trace.prompt_token_ids),
193
+ "steps": [
194
+ {"chosen": _prob_to_list(step.chosen), "top": [_prob_to_list(p) for p in step.top]}
195
+ for step in trace.steps
196
+ ],
197
+ }
198
+
199
+
200
+ def _trace_from_dict(data: dict[str, JSONValue]) -> ReferenceTrace:
201
+ ids = data["prompt_token_ids"]
202
+ return ReferenceTrace(
203
+ prompt_id=str(data["prompt_id"]),
204
+ prompt_token_ids=None if ids is None else tuple(int(i) for i in ids),
205
+ steps=tuple(
206
+ TokenStep(
207
+ chosen=_prob_from_list(step["chosen"]),
208
+ top=tuple(_prob_from_list(p) for p in step["top"]),
209
+ )
210
+ for step in data["steps"]
211
+ ),
212
+ )
213
+
214
+
215
+ def _prob_to_list(prob: TokenProb) -> list[JSONValue]:
216
+ """[token, logprob, token id or null, token bytes as base64 or null]."""
217
+ raw = None if prob.token_bytes is None else base64.b64encode(prob.token_bytes).decode("ascii")
218
+ return [prob.token, prob.logprob, prob.token_id, raw]
219
+
220
+
221
+ def _prob_from_list(data: JSONValue) -> TokenProb:
222
+ token, logprob, token_id, raw = data
223
+ return TokenProb(
224
+ token=str(token),
225
+ logprob=float(logprob),
226
+ token_id=None if token_id is None else int(token_id),
227
+ token_bytes=None if raw is None else base64.b64decode(str(raw), validate=True),
228
+ )
229
+
230
+
231
+ def _optional_str(value: JSONValue) -> str | None:
232
+ return None if value is None else str(value)
233
+
234
+
235
+ def _optional_int(value: JSONValue) -> int | None:
236
+ return None if value is None else int(value)
237
+
238
+
239
+ def _optional_float(value: JSONValue) -> float | None:
240
+ return None if value is None else float(value)