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,359 @@
|
|
|
1
|
+
"""Adapter for the Ollama native API (/api/generate, /api/chat, /api/show).
|
|
2
|
+
|
|
3
|
+
Ollama reports logprobs as tokens without ids, so teacher forcing works by text: each
|
|
4
|
+
scoring request sends the prompt plus the chosen tokens so far, rebuilt from the bytes
|
|
5
|
+
Ollama reports for every token. Ollama retokenizes that text, which can split it
|
|
6
|
+
differently from the reference's token sequence (for example around whitespace). That is
|
|
7
|
+
why `exact_token_ids` is False for this backend.
|
|
8
|
+
|
|
9
|
+
Byte-level tokenizers split many non-Latin characters across tokens. Ollama (checked on
|
|
10
|
+
0.35.1) holds back a token that ends inside a UTF-8 character:
|
|
11
|
+
|
|
12
|
+
- While generating, it reports the held-back tokens and the token that completes the
|
|
13
|
+
character as one entry: the text and bytes of the whole character, but the logprob and
|
|
14
|
+
alternatives of the last token only. Those alternatives are character fragments, all
|
|
15
|
+
shown as U+FFFD with the bytes of U+FFFD. That distribution belongs to a later position
|
|
16
|
+
than the one the entry stands for, and no text prompt can reproduce it, because a
|
|
17
|
+
prompt cannot end inside a character. `generate_scored` therefore keeps such a step but
|
|
18
|
+
empties its `top`, which marks the position as unscored for every candidate; metrics
|
|
19
|
+
skip it, so it counts neither as a top-1 miss nor in the KL divergence. The steps after
|
|
20
|
+
it are scored normally.
|
|
21
|
+
- When the one token a scoring request asks for ends inside a character, Ollama returns
|
|
22
|
+
no text and no logprobs, exactly as for end of sequence (only `done_reason` differs).
|
|
23
|
+
Either way no distribution is available, so the position is reported as an empty
|
|
24
|
+
TopK: a top-1 miss without a KL value. Against a reference from Ollama that is right,
|
|
25
|
+
because a scored reference step is always a whole token and the candidate picked a
|
|
26
|
+
fragment.
|
|
27
|
+
|
|
28
|
+
An empty TopK from `score_continuation` also stands for a position the reference left
|
|
29
|
+
unscored; no request is sent for it. Ollama leaves out the `logprobs` key entirely when
|
|
30
|
+
the first token it picks is end of sequence (or a held-back fragment), so a reply with no
|
|
31
|
+
text and no `logprobs` is read as zero steps rather than as missing logprob support.
|
|
32
|
+
|
|
33
|
+
quantdiff never sets `num_ctx`: it measures the context the user actually gets.
|
|
34
|
+
"""
|
|
35
|
+
|
|
36
|
+
from __future__ import annotations
|
|
37
|
+
|
|
38
|
+
import logging
|
|
39
|
+
import time
|
|
40
|
+
from collections.abc import Sequence
|
|
41
|
+
from typing import Final
|
|
42
|
+
|
|
43
|
+
from quantdiff._http import DEFAULT_TIMEOUT_SECONDS, get_json, post_json, validate_base_url
|
|
44
|
+
from quantdiff._text import printable
|
|
45
|
+
from quantdiff.backends._common import (
|
|
46
|
+
GREEDY_TEMPERATURE,
|
|
47
|
+
SCORING_SEED,
|
|
48
|
+
RequestFailure,
|
|
49
|
+
check_max_tokens,
|
|
50
|
+
check_top_k,
|
|
51
|
+
decode_rate,
|
|
52
|
+
expect_dict,
|
|
53
|
+
expect_float,
|
|
54
|
+
expect_list,
|
|
55
|
+
expect_str,
|
|
56
|
+
explained_failures,
|
|
57
|
+
forced_texts,
|
|
58
|
+
mentions_missing_model,
|
|
59
|
+
openai_messages,
|
|
60
|
+
openai_tool,
|
|
61
|
+
optional_bytes,
|
|
62
|
+
optional_int,
|
|
63
|
+
optional_str,
|
|
64
|
+
sorted_top,
|
|
65
|
+
tool_call,
|
|
66
|
+
unfolded_steps,
|
|
67
|
+
)
|
|
68
|
+
from quantdiff.errors import CapabilityError, SpecError
|
|
69
|
+
from quantdiff.types import (
|
|
70
|
+
CandidateSpec,
|
|
71
|
+
ChatResult,
|
|
72
|
+
JSONValue,
|
|
73
|
+
Message,
|
|
74
|
+
ServerInfo,
|
|
75
|
+
TokenProb,
|
|
76
|
+
TokenStep,
|
|
77
|
+
ToolCall,
|
|
78
|
+
ToolSpec,
|
|
79
|
+
TopK,
|
|
80
|
+
)
|
|
81
|
+
|
|
82
|
+
logger = logging.getLogger(__name__)
|
|
83
|
+
|
|
84
|
+
KEEP_ALIVE: Final = "10m"
|
|
85
|
+
"""Keeps the model loaded between the many small scoring requests."""
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
class OllamaBackend:
|
|
89
|
+
"""A model served by Ollama, addressed by its tag."""
|
|
90
|
+
|
|
91
|
+
def __init__(self, spec: CandidateSpec, *, timeout: float = DEFAULT_TIMEOUT_SECONDS) -> None:
|
|
92
|
+
if spec.kind != "ollama":
|
|
93
|
+
raise SpecError(f"OllamaBackend cannot serve a {spec.kind!r} spec")
|
|
94
|
+
if not spec.model:
|
|
95
|
+
raise SpecError("an Ollama spec needs a model tag")
|
|
96
|
+
self._spec = spec
|
|
97
|
+
self._base_url = validate_base_url(spec.base_url)
|
|
98
|
+
self._timeout = timeout
|
|
99
|
+
self._info: ServerInfo | None = None
|
|
100
|
+
|
|
101
|
+
@property
|
|
102
|
+
def spec(self) -> CandidateSpec:
|
|
103
|
+
return self._spec
|
|
104
|
+
|
|
105
|
+
def info(self) -> ServerInfo:
|
|
106
|
+
"""Describe the model.
|
|
107
|
+
|
|
108
|
+
The effective context length comes from a `num_ctx` set in the Modelfile, or else
|
|
109
|
+
from /api/ps after loading the model, because recent Ollama versions choose the
|
|
110
|
+
default context size at load time from available memory.
|
|
111
|
+
"""
|
|
112
|
+
if self._info is None:
|
|
113
|
+
self._info = self._load_info()
|
|
114
|
+
return self._info
|
|
115
|
+
|
|
116
|
+
def chat(
|
|
117
|
+
self,
|
|
118
|
+
messages: Sequence[Message],
|
|
119
|
+
*,
|
|
120
|
+
max_tokens: int,
|
|
121
|
+
tools: Sequence[ToolSpec] = (),
|
|
122
|
+
json_schema: dict[str, JSONValue] | None = None,
|
|
123
|
+
seed: int = 0,
|
|
124
|
+
) -> ChatResult:
|
|
125
|
+
check_max_tokens(max_tokens)
|
|
126
|
+
payload: dict[str, JSONValue] = {
|
|
127
|
+
"model": self._spec.model,
|
|
128
|
+
"messages": openai_messages(messages),
|
|
129
|
+
"stream": False,
|
|
130
|
+
"keep_alive": KEEP_ALIVE,
|
|
131
|
+
"options": {
|
|
132
|
+
"temperature": GREEDY_TEMPERATURE,
|
|
133
|
+
"seed": seed,
|
|
134
|
+
"num_predict": max_tokens,
|
|
135
|
+
},
|
|
136
|
+
}
|
|
137
|
+
if tools:
|
|
138
|
+
payload["tools"] = [openai_tool(tool) for tool in tools]
|
|
139
|
+
if json_schema is not None:
|
|
140
|
+
payload["format"] = json_schema
|
|
141
|
+
started = time.perf_counter()
|
|
142
|
+
body = self._post("/api/chat", payload)
|
|
143
|
+
seconds = time.perf_counter() - started
|
|
144
|
+
return _parse_chat(body, seconds=seconds)
|
|
145
|
+
|
|
146
|
+
def tokenize(self, text: str) -> tuple[int, ...] | None:
|
|
147
|
+
return None
|
|
148
|
+
|
|
149
|
+
def generate_scored(self, prompt: str, *, max_tokens: int, top_k: int) -> list[TokenStep]:
|
|
150
|
+
check_max_tokens(max_tokens)
|
|
151
|
+
check_top_k(top_k)
|
|
152
|
+
return unfolded_steps(self._generate(prompt, max_tokens=max_tokens, top_k=top_k))
|
|
153
|
+
|
|
154
|
+
def score_continuation(
|
|
155
|
+
self,
|
|
156
|
+
prompt: str,
|
|
157
|
+
continuation: Sequence[TokenStep],
|
|
158
|
+
*,
|
|
159
|
+
top_k: int,
|
|
160
|
+
prompt_token_ids: Sequence[int] | None = None,
|
|
161
|
+
) -> list[TopK]:
|
|
162
|
+
"""Teacher-force by text. `prompt_token_ids` is ignored; see the module docstring."""
|
|
163
|
+
check_top_k(top_k)
|
|
164
|
+
distributions: list[TopK] = []
|
|
165
|
+
for text in forced_texts(prompt, continuation):
|
|
166
|
+
steps = [] if text is None else self._generate(text, max_tokens=1, top_k=top_k)
|
|
167
|
+
distributions.append(steps[0].top if steps else ())
|
|
168
|
+
return distributions
|
|
169
|
+
|
|
170
|
+
def close(self) -> None:
|
|
171
|
+
"""Nothing to release: every request uses its own connection."""
|
|
172
|
+
|
|
173
|
+
def _generate(self, prompt: str, *, max_tokens: int, top_k: int) -> list[TokenStep]:
|
|
174
|
+
payload: dict[str, JSONValue] = {
|
|
175
|
+
"model": self._spec.model,
|
|
176
|
+
"prompt": prompt,
|
|
177
|
+
"raw": True,
|
|
178
|
+
"stream": False,
|
|
179
|
+
"logprobs": True,
|
|
180
|
+
"top_logprobs": top_k,
|
|
181
|
+
"keep_alive": KEEP_ALIVE,
|
|
182
|
+
"options": {
|
|
183
|
+
"num_predict": max_tokens,
|
|
184
|
+
"temperature": GREEDY_TEMPERATURE,
|
|
185
|
+
"seed": SCORING_SEED,
|
|
186
|
+
},
|
|
187
|
+
}
|
|
188
|
+
body = self._post("/api/generate", payload)
|
|
189
|
+
if "logprobs" not in body:
|
|
190
|
+
if _stopped_without_text(body):
|
|
191
|
+
return []
|
|
192
|
+
raise CapabilityError(
|
|
193
|
+
f"Ollama at {self._base_url} returned no logprobs; upgrade to Ollama 0.12 or newer"
|
|
194
|
+
)
|
|
195
|
+
entries = body["logprobs"] or []
|
|
196
|
+
return [_parse_step(entry) for entry in expect_list(entries, "Ollama logprobs")]
|
|
197
|
+
|
|
198
|
+
def _load_info(self) -> ServerInfo:
|
|
199
|
+
show = self._post("/api/show", {"model": self._spec.model})
|
|
200
|
+
version = self._get("/api/version")
|
|
201
|
+
details = expect_dict(show.get("details") or {}, "Ollama model details")
|
|
202
|
+
listed = self._listed_entry("/api/tags") or {}
|
|
203
|
+
digest = optional_str(listed.get("digest"))
|
|
204
|
+
facts = {
|
|
205
|
+
"quantization": optional_str(details.get("quantization_level")),
|
|
206
|
+
"parameter_size": optional_str(details.get("parameter_size")),
|
|
207
|
+
"trained_context_length": _optional_text(_trained_context(show.get("model_info"))),
|
|
208
|
+
"server_version": optional_str(version.get("version")),
|
|
209
|
+
"digest": digest,
|
|
210
|
+
# Only a fallback fingerprint: it changes when a tag is re-pulled, like the digest.
|
|
211
|
+
"modified_at": None if digest else optional_str(show.get("modified_at")),
|
|
212
|
+
}
|
|
213
|
+
context_length = _num_ctx_parameter(show.get("parameters"))
|
|
214
|
+
if context_length is None:
|
|
215
|
+
context_length = self._loaded_context_length()
|
|
216
|
+
return ServerInfo(
|
|
217
|
+
backend="ollama",
|
|
218
|
+
model=self._spec.model,
|
|
219
|
+
context_length=context_length,
|
|
220
|
+
chat_template=optional_str(show.get("template")),
|
|
221
|
+
template_dialect="go",
|
|
222
|
+
supports_logprobs=True,
|
|
223
|
+
exact_token_ids=False,
|
|
224
|
+
details=tuple((key, value) for key, value in facts.items() if value),
|
|
225
|
+
size_bytes=optional_int(listed.get("size")),
|
|
226
|
+
weights_id=digest,
|
|
227
|
+
)
|
|
228
|
+
|
|
229
|
+
def _loaded_context_length(self) -> int | None:
|
|
230
|
+
# A generate request without a prompt only loads the model, after which /api/ps
|
|
231
|
+
# reports the context size Ollama actually allocated.
|
|
232
|
+
self._post("/api/generate", {"model": self._spec.model, "keep_alive": KEEP_ALIVE})
|
|
233
|
+
model = self._listed_entry("/api/ps")
|
|
234
|
+
if model is None:
|
|
235
|
+
logger.debug("model %s is not listed by /api/ps", self._spec.model)
|
|
236
|
+
return None
|
|
237
|
+
return optional_int(model.get("context_length"))
|
|
238
|
+
|
|
239
|
+
def _listed_entry(self, path: str) -> dict[str, JSONValue] | None:
|
|
240
|
+
"""Find this model in a /api/ps or /api/tags listing."""
|
|
241
|
+
listing = self._get(path)
|
|
242
|
+
names = _tag_aliases(self._spec.model)
|
|
243
|
+
for entry in expect_list(listing.get("models") or [], f"Ollama {path} models"):
|
|
244
|
+
model = expect_dict(entry, f"Ollama {path} entry")
|
|
245
|
+
if model.get("name") in names or model.get("model") in names:
|
|
246
|
+
return model
|
|
247
|
+
return None
|
|
248
|
+
|
|
249
|
+
def _get(self, path: str) -> dict[str, JSONValue]:
|
|
250
|
+
with explained_failures(self._explain):
|
|
251
|
+
body = get_json(self._url(path), timeout=self._timeout)
|
|
252
|
+
return expect_dict(body, f"Ollama {path} response")
|
|
253
|
+
|
|
254
|
+
def _post(self, path: str, payload: dict[str, JSONValue]) -> dict[str, JSONValue]:
|
|
255
|
+
with explained_failures(self._explain):
|
|
256
|
+
body = post_json(self._url(path), payload, timeout=self._timeout)
|
|
257
|
+
return expect_dict(body, f"Ollama {path} response")
|
|
258
|
+
|
|
259
|
+
def _explain(self, failure: RequestFailure) -> str | None:
|
|
260
|
+
return explain_ollama_failure(failure, base_url=self._base_url, model=self._spec.model)
|
|
261
|
+
|
|
262
|
+
def _url(self, path: str) -> str:
|
|
263
|
+
return self._base_url + path
|
|
264
|
+
|
|
265
|
+
|
|
266
|
+
def explain_ollama_failure(
|
|
267
|
+
failure: RequestFailure, *, base_url: str, model: str | None = None
|
|
268
|
+
) -> str | None:
|
|
269
|
+
"""Say what to do about a failed Ollama request, or return None if there is no advice."""
|
|
270
|
+
if failure.status is None:
|
|
271
|
+
return (
|
|
272
|
+
f"cannot reach Ollama at {base_url} ({failure.detail}); "
|
|
273
|
+
"start it with `ollama serve` or set OLLAMA_HOST"
|
|
274
|
+
)
|
|
275
|
+
if model and failure.status == 404 and mentions_missing_model(failure.detail):
|
|
276
|
+
name = printable(model)
|
|
277
|
+
return (
|
|
278
|
+
f"model {name!r} is not available in Ollama; "
|
|
279
|
+
f"run `ollama pull {name}` (see `ollama list`)"
|
|
280
|
+
)
|
|
281
|
+
return None
|
|
282
|
+
|
|
283
|
+
|
|
284
|
+
def _stopped_without_text(body: dict[str, JSONValue]) -> bool:
|
|
285
|
+
return body.get("done") is True and not body.get("response")
|
|
286
|
+
|
|
287
|
+
|
|
288
|
+
def _parse_chat(body: dict[str, JSONValue], *, seconds: float) -> ChatResult:
|
|
289
|
+
message = expect_dict(body.get("message"), "Ollama chat message")
|
|
290
|
+
raw_calls = expect_list(message.get("tool_calls") or [], "Ollama tool_calls")
|
|
291
|
+
return ChatResult(
|
|
292
|
+
text=optional_str(message.get("content")) or "",
|
|
293
|
+
tool_calls=tuple(_parse_tool_call(call) for call in raw_calls),
|
|
294
|
+
finish_reason=optional_str(body.get("done_reason")),
|
|
295
|
+
prompt_tokens=optional_int(body.get("prompt_eval_count")),
|
|
296
|
+
completion_tokens=optional_int(body.get("eval_count")),
|
|
297
|
+
seconds=seconds,
|
|
298
|
+
decode_tokens_per_second=decode_rate(
|
|
299
|
+
optional_int(body.get("eval_count")), _nanoseconds(body.get("eval_duration"))
|
|
300
|
+
),
|
|
301
|
+
)
|
|
302
|
+
|
|
303
|
+
|
|
304
|
+
def _nanoseconds(value: JSONValue) -> float | None:
|
|
305
|
+
"""Convert one of Ollama's integer nanosecond durations to seconds."""
|
|
306
|
+
duration = optional_int(value)
|
|
307
|
+
return None if duration is None else duration / 1e9
|
|
308
|
+
|
|
309
|
+
|
|
310
|
+
def _parse_tool_call(call: JSONValue) -> ToolCall:
|
|
311
|
+
function = expect_dict(expect_dict(call, "Ollama tool call").get("function"), "function")
|
|
312
|
+
return tool_call(expect_str(function.get("name"), "tool call name"), function.get("arguments"))
|
|
313
|
+
|
|
314
|
+
|
|
315
|
+
def _parse_step(entry: JSONValue) -> TokenStep:
|
|
316
|
+
step = expect_dict(entry, "Ollama logprob entry")
|
|
317
|
+
top = expect_list(step.get("top_logprobs") or [], "Ollama top_logprobs")
|
|
318
|
+
return TokenStep(chosen=_parse_prob(step), top=sorted_top(_parse_prob(item) for item in top))
|
|
319
|
+
|
|
320
|
+
|
|
321
|
+
def _parse_prob(value: JSONValue) -> TokenProb:
|
|
322
|
+
item = expect_dict(value, "Ollama logprob")
|
|
323
|
+
return TokenProb(
|
|
324
|
+
token=expect_str(item.get("token"), "Ollama token"),
|
|
325
|
+
logprob=expect_float(item.get("logprob"), "Ollama logprob"),
|
|
326
|
+
token_bytes=optional_bytes(item.get("bytes")),
|
|
327
|
+
)
|
|
328
|
+
|
|
329
|
+
|
|
330
|
+
def _num_ctx_parameter(parameters: JSONValue) -> int | None:
|
|
331
|
+
"""Read `num_ctx` from the Modelfile parameter block, one `name value` pair per line."""
|
|
332
|
+
if not isinstance(parameters, str):
|
|
333
|
+
return None
|
|
334
|
+
for line in parameters.splitlines():
|
|
335
|
+
fields = line.split()
|
|
336
|
+
if len(fields) == 2 and fields[0] == "num_ctx" and fields[1].isdigit():
|
|
337
|
+
return int(fields[1])
|
|
338
|
+
return None
|
|
339
|
+
|
|
340
|
+
|
|
341
|
+
def _trained_context(model_info: JSONValue) -> int | None:
|
|
342
|
+
if not isinstance(model_info, dict):
|
|
343
|
+
return None
|
|
344
|
+
for key, value in model_info.items():
|
|
345
|
+
if key.endswith(".context_length"):
|
|
346
|
+
return optional_int(value)
|
|
347
|
+
return None
|
|
348
|
+
|
|
349
|
+
|
|
350
|
+
def _optional_text(value: int | None) -> str | None:
|
|
351
|
+
return None if value is None else str(value)
|
|
352
|
+
|
|
353
|
+
|
|
354
|
+
def _tag_aliases(tag: str) -> frozenset[str]:
|
|
355
|
+
"""Names /api/ps may use for `tag`; Ollama adds ':latest' when no tag is given."""
|
|
356
|
+
_, _, name = tag.rpartition("/")
|
|
357
|
+
if ":" in name:
|
|
358
|
+
return frozenset({tag})
|
|
359
|
+
return frozenset({tag, f"{tag}:latest"})
|
|
@@ -0,0 +1,338 @@
|
|
|
1
|
+
"""Adapter for OpenAI-compatible servers such as LM Studio and vLLM.
|
|
2
|
+
|
|
3
|
+
Scoring uses the legacy /completions endpoint. Its `logprobs` field holds either the
|
|
4
|
+
legacy parallel `tokens`, `token_logprobs` and `top_logprobs` arrays, or (llama-server)
|
|
5
|
+
a `content` list in the chat format, whose entries also carry each token's bytes. Tokens
|
|
6
|
+
are read as text without ids, so teacher forcing works by text and the server
|
|
7
|
+
retokenizes each prefix; see `exact_token_ids` in ServerInfo. Prefixes are rebuilt from
|
|
8
|
+
token bytes when the server reported them.
|
|
9
|
+
|
|
10
|
+
Servers that hold back a token ending inside a UTF-8 character (llama-server does) fold
|
|
11
|
+
it into the entry of the token that completes the character, an entry whose
|
|
12
|
+
alternatives belong to the last folded token. `generate_scored` keeps such a step but
|
|
13
|
+
empties its `top`, which marks the position as unscored for every candidate. A
|
|
14
|
+
one-token request whose token is such a fragment comes back with the text U+FFFD and no
|
|
15
|
+
logprobs; it is returned as an empty TopK, a top-1 miss without a KL value.
|
|
16
|
+
|
|
17
|
+
An empty TopK from `score_continuation` otherwise means the server ended generation at
|
|
18
|
+
that position (end of sequence) without reporting a distribution, or the reference left
|
|
19
|
+
the position unscored, in which case no request is sent.
|
|
20
|
+
|
|
21
|
+
The legacy `top_logprobs` entry is an object keyed by token text, so two distinct tokens
|
|
22
|
+
that decode to the same text (for example a byte-fallback token and a merged token)
|
|
23
|
+
arrive as one entry and only one of their logprobs survives.
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
from __future__ import annotations
|
|
27
|
+
|
|
28
|
+
import os
|
|
29
|
+
import time
|
|
30
|
+
from collections.abc import Mapping, Sequence
|
|
31
|
+
from typing import Final
|
|
32
|
+
|
|
33
|
+
from quantdiff._http import DEFAULT_TIMEOUT_SECONDS, get_json, post_json, validate_base_url
|
|
34
|
+
from quantdiff._text import printable
|
|
35
|
+
from quantdiff.backends._common import (
|
|
36
|
+
GREEDY_TEMPERATURE,
|
|
37
|
+
SCORING_SEED,
|
|
38
|
+
RequestFailure,
|
|
39
|
+
check_max_tokens,
|
|
40
|
+
check_top_k,
|
|
41
|
+
expect_dict,
|
|
42
|
+
expect_float,
|
|
43
|
+
expect_list,
|
|
44
|
+
expect_str,
|
|
45
|
+
explained_failures,
|
|
46
|
+
forced_texts,
|
|
47
|
+
mentions_missing_model,
|
|
48
|
+
openai_chat_payload,
|
|
49
|
+
optional_bytes,
|
|
50
|
+
optional_int,
|
|
51
|
+
parse_openai_chat,
|
|
52
|
+
sorted_top,
|
|
53
|
+
unfolded_steps,
|
|
54
|
+
)
|
|
55
|
+
from quantdiff.errors import BackendError, CapabilityError, SpecError
|
|
56
|
+
from quantdiff.types import (
|
|
57
|
+
CandidateSpec,
|
|
58
|
+
ChatResult,
|
|
59
|
+
JSONValue,
|
|
60
|
+
Message,
|
|
61
|
+
ServerInfo,
|
|
62
|
+
TokenProb,
|
|
63
|
+
TokenStep,
|
|
64
|
+
ToolSpec,
|
|
65
|
+
TopK,
|
|
66
|
+
)
|
|
67
|
+
|
|
68
|
+
_CONTEXT_FIELDS: Final = ("max_model_len", "loaded_context_length", "context_length")
|
|
69
|
+
"""Model-list fields that report context size: vLLM first, then LM Studio."""
|
|
70
|
+
_LISTED_MODELS_IN_ERROR: Final = 10
|
|
71
|
+
_AUTH_STATUSES: Final = frozenset({401, 403})
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
class OpenAICompatBackend:
|
|
75
|
+
"""A named model on a server that speaks the OpenAI /v1 API."""
|
|
76
|
+
|
|
77
|
+
def __init__(self, spec: CandidateSpec, *, timeout: float = DEFAULT_TIMEOUT_SECONDS) -> None:
|
|
78
|
+
if spec.kind != "openai":
|
|
79
|
+
raise SpecError(f"OpenAICompatBackend cannot serve a {spec.kind!r} spec")
|
|
80
|
+
if not spec.model:
|
|
81
|
+
raise SpecError("an OpenAI-compatible spec needs a model name")
|
|
82
|
+
self._spec = spec
|
|
83
|
+
self._base_url = validate_base_url(spec.base_url)
|
|
84
|
+
self._timeout = timeout
|
|
85
|
+
self._headers = _auth_headers(spec.api_key_env)
|
|
86
|
+
self._info: ServerInfo | None = None
|
|
87
|
+
|
|
88
|
+
@property
|
|
89
|
+
def spec(self) -> CandidateSpec:
|
|
90
|
+
return self._spec
|
|
91
|
+
|
|
92
|
+
def info(self) -> ServerInfo:
|
|
93
|
+
"""Describe the model. Logprob support is assumed until a scored call proves otherwise."""
|
|
94
|
+
if self._info is None:
|
|
95
|
+
self._info = self._load_info()
|
|
96
|
+
return self._info
|
|
97
|
+
|
|
98
|
+
def chat(
|
|
99
|
+
self,
|
|
100
|
+
messages: Sequence[Message],
|
|
101
|
+
*,
|
|
102
|
+
max_tokens: int,
|
|
103
|
+
tools: Sequence[ToolSpec] = (),
|
|
104
|
+
json_schema: dict[str, JSONValue] | None = None,
|
|
105
|
+
seed: int = 0,
|
|
106
|
+
) -> ChatResult:
|
|
107
|
+
payload = openai_chat_payload(
|
|
108
|
+
self._spec.model,
|
|
109
|
+
messages,
|
|
110
|
+
max_tokens=max_tokens,
|
|
111
|
+
tools=tools,
|
|
112
|
+
json_schema=json_schema,
|
|
113
|
+
seed=seed,
|
|
114
|
+
)
|
|
115
|
+
started = time.perf_counter()
|
|
116
|
+
body = self._post("/chat/completions", payload)
|
|
117
|
+
return parse_openai_chat(body, seconds=time.perf_counter() - started)
|
|
118
|
+
|
|
119
|
+
def tokenize(self, text: str) -> tuple[int, ...] | None:
|
|
120
|
+
return None
|
|
121
|
+
|
|
122
|
+
def generate_scored(self, prompt: str, *, max_tokens: int, top_k: int) -> list[TokenStep]:
|
|
123
|
+
check_max_tokens(max_tokens)
|
|
124
|
+
check_top_k(top_k)
|
|
125
|
+
return unfolded_steps(self._complete(prompt, max_tokens=max_tokens, top_k=top_k))
|
|
126
|
+
|
|
127
|
+
def score_continuation(
|
|
128
|
+
self,
|
|
129
|
+
prompt: str,
|
|
130
|
+
continuation: Sequence[TokenStep],
|
|
131
|
+
*,
|
|
132
|
+
top_k: int,
|
|
133
|
+
prompt_token_ids: Sequence[int] | None = None,
|
|
134
|
+
) -> list[TopK]:
|
|
135
|
+
"""Teacher-force by text. `prompt_token_ids` is ignored; see the module docstring."""
|
|
136
|
+
check_top_k(top_k)
|
|
137
|
+
distributions: list[TopK] = []
|
|
138
|
+
for text in forced_texts(prompt, continuation):
|
|
139
|
+
steps = [] if text is None else self._complete(text, max_tokens=1, top_k=top_k)
|
|
140
|
+
distributions.append(steps[0].top if steps else ())
|
|
141
|
+
return distributions
|
|
142
|
+
|
|
143
|
+
def close(self) -> None:
|
|
144
|
+
"""Nothing to release: every request uses its own connection."""
|
|
145
|
+
|
|
146
|
+
def _complete(self, prompt: str, *, max_tokens: int, top_k: int) -> list[TokenStep]:
|
|
147
|
+
payload: dict[str, JSONValue] = {
|
|
148
|
+
"model": self._spec.model,
|
|
149
|
+
"prompt": prompt,
|
|
150
|
+
"max_tokens": max_tokens,
|
|
151
|
+
"temperature": GREEDY_TEMPERATURE,
|
|
152
|
+
"seed": SCORING_SEED,
|
|
153
|
+
"logprobs": top_k,
|
|
154
|
+
"stream": False,
|
|
155
|
+
}
|
|
156
|
+
body = self._post("/completions", payload)
|
|
157
|
+
choices = expect_list(body.get("choices"), "completions choices")
|
|
158
|
+
if not choices:
|
|
159
|
+
raise BackendError("completions response has no choices")
|
|
160
|
+
choice = expect_dict(choices[0], "completions choice")
|
|
161
|
+
logprobs = choice.get("logprobs")
|
|
162
|
+
if isinstance(logprobs, dict):
|
|
163
|
+
return _parse_logprobs(logprobs)
|
|
164
|
+
if max_tokens == 1 and _held_back(choice):
|
|
165
|
+
return []
|
|
166
|
+
raise CapabilityError(
|
|
167
|
+
f"{self._base_url} returned no logprobs for /completions; "
|
|
168
|
+
"this server cannot be used for logit metrics"
|
|
169
|
+
)
|
|
170
|
+
|
|
171
|
+
def _load_info(self) -> ServerInfo:
|
|
172
|
+
with explained_failures(self._explain):
|
|
173
|
+
listing = get_json(self._url("/models"), headers=self._headers, timeout=self._timeout)
|
|
174
|
+
body = expect_dict(listing, "model list")
|
|
175
|
+
entries = [
|
|
176
|
+
expect_dict(entry, "model list entry")
|
|
177
|
+
for entry in expect_list(body.get("data"), "model list data")
|
|
178
|
+
]
|
|
179
|
+
model = self._find_model(entries)
|
|
180
|
+
return ServerInfo(
|
|
181
|
+
backend="openai",
|
|
182
|
+
model=self._spec.model,
|
|
183
|
+
context_length=_context_length(model),
|
|
184
|
+
chat_template=None,
|
|
185
|
+
template_dialect="unknown",
|
|
186
|
+
supports_logprobs=True,
|
|
187
|
+
exact_token_ids=False,
|
|
188
|
+
details=_details(model),
|
|
189
|
+
)
|
|
190
|
+
|
|
191
|
+
def _find_model(self, entries: list[dict[str, JSONValue]]) -> dict[str, JSONValue]:
|
|
192
|
+
for entry in entries:
|
|
193
|
+
if entry.get("id") == self._spec.model:
|
|
194
|
+
return entry
|
|
195
|
+
if not entries:
|
|
196
|
+
raise BackendError(
|
|
197
|
+
f"model {self._spec.model!r} is not served by {self._base_url}, which lists no "
|
|
198
|
+
"models; load one in LM Studio or start vLLM with --model"
|
|
199
|
+
)
|
|
200
|
+
raise BackendError(
|
|
201
|
+
f"model {self._spec.model!r} is not served by {self._base_url}; "
|
|
202
|
+
f"available: {_model_ids(entries)}. Put one of these after '#' in the spec"
|
|
203
|
+
)
|
|
204
|
+
|
|
205
|
+
def _post(self, path: str, payload: dict[str, JSONValue]) -> dict[str, JSONValue]:
|
|
206
|
+
with explained_failures(self._explain):
|
|
207
|
+
body = post_json(self._url(path), payload, headers=self._headers, timeout=self._timeout)
|
|
208
|
+
return expect_dict(body, f"{path} response")
|
|
209
|
+
|
|
210
|
+
def _explain(self, failure: RequestFailure) -> str | None:
|
|
211
|
+
if failure.status is None:
|
|
212
|
+
return (
|
|
213
|
+
f"cannot reach the server at {self._base_url} ({failure.detail}); "
|
|
214
|
+
"check that LM Studio/vLLM is running and the URL ends in /v1"
|
|
215
|
+
)
|
|
216
|
+
if failure.status in _AUTH_STATUSES:
|
|
217
|
+
return self._auth_advice(failure.status)
|
|
218
|
+
if failure.status != 404:
|
|
219
|
+
return None
|
|
220
|
+
if mentions_missing_model(failure.detail):
|
|
221
|
+
return (
|
|
222
|
+
f"model {self._spec.model!r} is not available on {self._base_url}; "
|
|
223
|
+
f"use a model id listed at {self._base_url}/models"
|
|
224
|
+
)
|
|
225
|
+
if not self._base_url.endswith("/v1"):
|
|
226
|
+
return (
|
|
227
|
+
f"{printable(failure.url)} returned HTTP 404; "
|
|
228
|
+
"OpenAI-compatible URLs usually end in /v1, for example http://127.0.0.1:1234/v1"
|
|
229
|
+
)
|
|
230
|
+
return None
|
|
231
|
+
|
|
232
|
+
def _auth_advice(self, status: int) -> str:
|
|
233
|
+
if self._spec.api_key_env is None:
|
|
234
|
+
return (
|
|
235
|
+
f"the server at {self._base_url} requires an API key (HTTP {status}); put the key "
|
|
236
|
+
"in an environment variable and add @env:NAME to the spec"
|
|
237
|
+
)
|
|
238
|
+
return (
|
|
239
|
+
f"the server at {self._base_url} rejected the API key from "
|
|
240
|
+
f"${self._spec.api_key_env} (HTTP {status}); check the key and its permissions"
|
|
241
|
+
)
|
|
242
|
+
|
|
243
|
+
def _url(self, path: str) -> str:
|
|
244
|
+
return self._base_url + path
|
|
245
|
+
|
|
246
|
+
|
|
247
|
+
def _auth_headers(api_key_env: str | None) -> Mapping[str, str]:
|
|
248
|
+
if api_key_env is None:
|
|
249
|
+
return {}
|
|
250
|
+
key = os.environ.get(api_key_env)
|
|
251
|
+
if not key:
|
|
252
|
+
raise BackendError(f"environment variable {api_key_env} is not set or is empty")
|
|
253
|
+
return {"Authorization": f"Bearer {key}"}
|
|
254
|
+
|
|
255
|
+
|
|
256
|
+
def _model_ids(entries: list[dict[str, JSONValue]]) -> str:
|
|
257
|
+
"""Comma list of served ids, capped so a large hub does not flood the terminal."""
|
|
258
|
+
ids = [printable(str(entry.get("id"))) for entry in entries]
|
|
259
|
+
shown = ", ".join(ids[:_LISTED_MODELS_IN_ERROR])
|
|
260
|
+
hidden = len(ids) - _LISTED_MODELS_IN_ERROR
|
|
261
|
+
return f"{shown} and {hidden} more" if hidden > 0 else shown
|
|
262
|
+
|
|
263
|
+
|
|
264
|
+
def _context_length(model: dict[str, JSONValue]) -> int | None:
|
|
265
|
+
for field in _CONTEXT_FIELDS:
|
|
266
|
+
value = optional_int(model.get(field))
|
|
267
|
+
if value is not None:
|
|
268
|
+
return value
|
|
269
|
+
return None
|
|
270
|
+
|
|
271
|
+
|
|
272
|
+
def _details(model: dict[str, JSONValue]) -> tuple[tuple[str, str], ...]:
|
|
273
|
+
"""Owner and creation time, the only identity /models exposes for the reference cache."""
|
|
274
|
+
owner = model.get("owned_by")
|
|
275
|
+
created = optional_int(model.get("created"))
|
|
276
|
+
facts = {
|
|
277
|
+
"owned_by": owner if isinstance(owner, str) else None,
|
|
278
|
+
"created": None if created is None else str(created),
|
|
279
|
+
}
|
|
280
|
+
return tuple((key, value) for key, value in facts.items() if value)
|
|
281
|
+
|
|
282
|
+
|
|
283
|
+
def _held_back(choice: dict[str, JSONValue]) -> bool:
|
|
284
|
+
"""True when the server generated a fragment of a UTF-8 character, which it shows as
|
|
285
|
+
U+FFFD, and reported no logprobs for it."""
|
|
286
|
+
text = choice.get("text")
|
|
287
|
+
return isinstance(text, str) and text.endswith("\N{REPLACEMENT CHARACTER}")
|
|
288
|
+
|
|
289
|
+
|
|
290
|
+
def _parse_logprobs(logprobs: dict[str, JSONValue]) -> list[TokenStep]:
|
|
291
|
+
if "content" in logprobs:
|
|
292
|
+
entries = expect_list(logprobs["content"], "logprobs content")
|
|
293
|
+
return [_parse_content_step(entry) for entry in entries]
|
|
294
|
+
tokens = expect_list(logprobs.get("tokens"), "logprobs tokens")
|
|
295
|
+
chosen = expect_list(logprobs.get("token_logprobs"), "logprobs token_logprobs")
|
|
296
|
+
tops = expect_list(logprobs.get("top_logprobs"), "logprobs top_logprobs")
|
|
297
|
+
if not len(tokens) == len(chosen) == len(tops):
|
|
298
|
+
raise BackendError("logprobs arrays have different lengths")
|
|
299
|
+
return [
|
|
300
|
+
TokenStep(
|
|
301
|
+
chosen=TokenProb(
|
|
302
|
+
token=expect_str(token, "logprobs token"),
|
|
303
|
+
logprob=expect_float(logprob, "token logprob"),
|
|
304
|
+
),
|
|
305
|
+
top=_parse_top(top),
|
|
306
|
+
)
|
|
307
|
+
for token, logprob, top in zip(tokens, chosen, tops, strict=True)
|
|
308
|
+
]
|
|
309
|
+
|
|
310
|
+
|
|
311
|
+
def _parse_top(value: JSONValue) -> TopK:
|
|
312
|
+
alternatives = expect_dict(value, "top_logprobs entry")
|
|
313
|
+
return sorted_top(
|
|
314
|
+
TokenProb(token=token, logprob=expect_float(logprob, "top logprob"))
|
|
315
|
+
for token, logprob in alternatives.items()
|
|
316
|
+
)
|
|
317
|
+
|
|
318
|
+
|
|
319
|
+
def _parse_content_step(entry: JSONValue) -> TokenStep:
|
|
320
|
+
"""Read one entry of the chat-format `content` list. An alternative whose logprob is
|
|
321
|
+
null (minus infinity) has zero probability and is left out."""
|
|
322
|
+
step = expect_dict(entry, "logprobs content entry")
|
|
323
|
+
top = expect_list(step.get("top_logprobs") or [], "content top_logprobs")
|
|
324
|
+
alternatives = (_parse_content_prob(item) for item in top if not _is_impossible(item))
|
|
325
|
+
return TokenStep(chosen=_parse_content_prob(step), top=sorted_top(alternatives))
|
|
326
|
+
|
|
327
|
+
|
|
328
|
+
def _is_impossible(item: JSONValue) -> bool:
|
|
329
|
+
return isinstance(item, dict) and "logprob" in item and item["logprob"] is None
|
|
330
|
+
|
|
331
|
+
|
|
332
|
+
def _parse_content_prob(value: JSONValue) -> TokenProb:
|
|
333
|
+
item = expect_dict(value, "logprobs content token")
|
|
334
|
+
return TokenProb(
|
|
335
|
+
token=expect_str(item.get("token"), "logprobs content token"),
|
|
336
|
+
logprob=expect_float(item.get("logprob"), "token logprob"),
|
|
337
|
+
token_bytes=optional_bytes(item.get("bytes")),
|
|
338
|
+
)
|