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
@@ -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
+ )