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,428 @@
|
|
|
1
|
+
"""Adapter for llama.cpp's llama-server (/completion, /tokenize, /props, /v1/chat/completions).
|
|
2
|
+
|
|
3
|
+
llama-server accepts prompts as token id arrays and reports the id of every token it
|
|
4
|
+
scores, so teacher forcing feeds the reference's exact token sequence and never
|
|
5
|
+
retokenizes text. When the reference came from a backend without token ids (Ollama or an
|
|
6
|
+
OpenAI-compatible server), teacher forcing falls back to text prompts, which llama-server
|
|
7
|
+
tokenizes itself. Both the current `completion_probabilities` format (id, token,
|
|
8
|
+
logprob, top_logprobs) and the older one (content, probs with tok_str and prob) are read.
|
|
9
|
+
|
|
10
|
+
llama-server writes a logprob of minus infinity as JSON null. Such alternatives have zero
|
|
11
|
+
probability and are left out of the top-k list.
|
|
12
|
+
|
|
13
|
+
Byte-level tokenizers split many non-Latin characters across tokens, and llama-server
|
|
14
|
+
(checked on b11425) holds back a token that ends inside a UTF-8 character:
|
|
15
|
+
|
|
16
|
+
- In a multi-token completion it reports the held-back tokens and the token that
|
|
17
|
+
completes the character as one entry, with the text and bytes of the whole character
|
|
18
|
+
but the id, logprob and alternatives of the last token only. `generate_scored` keeps
|
|
19
|
+
the leading entries that stand for one token each and, from the first folded entry on,
|
|
20
|
+
generates one token per request, so the trace holds every token id and every position's
|
|
21
|
+
own distribution and exact teacher forcing feeds the true token sequence.
|
|
22
|
+
- When the one token a request asks for ends inside a character, the reply has no
|
|
23
|
+
`completion_probabilities` at all. The request is then repeated with a logit bias that
|
|
24
|
+
makes the server pick a newline token instead. Probabilities are reported from the raw
|
|
25
|
+
logits (the default, without `post_sampling_probs`), so the alternatives in that reply
|
|
26
|
+
are the unbiased distribution at the same position, and under greedy decoding its head
|
|
27
|
+
is the token the server held back. Only a server that reports no probabilities even
|
|
28
|
+
for that complete token raises CapabilityError.
|
|
29
|
+
|
|
30
|
+
When the reference came from a backend without token ids, a position whose text prefix
|
|
31
|
+
would end inside a character cannot be sent as a text prompt and is returned as an empty
|
|
32
|
+
TopK without a request, as are positions the reference left unscored.
|
|
33
|
+
"""
|
|
34
|
+
|
|
35
|
+
from __future__ import annotations
|
|
36
|
+
|
|
37
|
+
import time
|
|
38
|
+
import urllib.parse
|
|
39
|
+
from collections.abc import Sequence
|
|
40
|
+
from dataclasses import dataclass, replace
|
|
41
|
+
from itertools import takewhile
|
|
42
|
+
from pathlib import PureWindowsPath
|
|
43
|
+
from typing import Final
|
|
44
|
+
|
|
45
|
+
from quantdiff._http import DEFAULT_TIMEOUT_SECONDS, get_json, post_json, validate_base_url
|
|
46
|
+
from quantdiff._text import printable
|
|
47
|
+
from quantdiff.backends._common import (
|
|
48
|
+
GREEDY_TEMPERATURE,
|
|
49
|
+
SCORING_SEED,
|
|
50
|
+
RequestFailure,
|
|
51
|
+
check_max_tokens,
|
|
52
|
+
check_top_k,
|
|
53
|
+
decode_rate,
|
|
54
|
+
expect_dict,
|
|
55
|
+
expect_float,
|
|
56
|
+
expect_list,
|
|
57
|
+
expect_str,
|
|
58
|
+
explained_failures,
|
|
59
|
+
folds_tokens,
|
|
60
|
+
forced_texts,
|
|
61
|
+
logprob_from_prob,
|
|
62
|
+
openai_chat_payload,
|
|
63
|
+
optional_bytes,
|
|
64
|
+
optional_int,
|
|
65
|
+
optional_str,
|
|
66
|
+
parse_openai_chat,
|
|
67
|
+
sorted_top,
|
|
68
|
+
)
|
|
69
|
+
from quantdiff.errors import BackendError, CapabilityError, SpecError
|
|
70
|
+
from quantdiff.types import (
|
|
71
|
+
CandidateSpec,
|
|
72
|
+
ChatResult,
|
|
73
|
+
JSONValue,
|
|
74
|
+
Message,
|
|
75
|
+
ServerInfo,
|
|
76
|
+
TokenProb,
|
|
77
|
+
TokenStep,
|
|
78
|
+
ToolSpec,
|
|
79
|
+
TopK,
|
|
80
|
+
)
|
|
81
|
+
|
|
82
|
+
Prompt = str | tuple[int, ...]
|
|
83
|
+
"""Raw text, which the server tokenizes with special tokens added, or exact token ids."""
|
|
84
|
+
|
|
85
|
+
_GREEDY_SAMPLERS: Final = ("temperature",)
|
|
86
|
+
"""Only the temperature stage, so server-side penalties cannot override the argmax token."""
|
|
87
|
+
_DEFAULT_PORT: Final = 8080
|
|
88
|
+
_LOADING_STATUS: Final = 503
|
|
89
|
+
_META_FINGERPRINT: Final = ("size", "n_params", "n_vocab")
|
|
90
|
+
"""/v1/models meta fields that identify the loaded GGUF file."""
|
|
91
|
+
_STOP_TYPES: Final = frozenset({"eos", "word"})
|
|
92
|
+
"""`stop_type` values for a completion the server ended itself, before the token limit."""
|
|
93
|
+
_COMPLETE_TEXT: Final = "\n"
|
|
94
|
+
"""An ASCII character, so its token is complete text. It is forced when the token the
|
|
95
|
+
server picked ends inside a character."""
|
|
96
|
+
_FORCING_BIAS: Final = 1000.0
|
|
97
|
+
"""Logit bias that makes greedy decoding pick the forced token."""
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
@dataclass(frozen=True, slots=True)
|
|
101
|
+
class _Completion:
|
|
102
|
+
steps: tuple[TokenStep, ...]
|
|
103
|
+
reported: bool
|
|
104
|
+
"""True when the reply carried `completion_probabilities`, even an empty list."""
|
|
105
|
+
stopped: bool
|
|
106
|
+
"""True when the server ended generation itself (end of sequence or a stop word)."""
|
|
107
|
+
|
|
108
|
+
|
|
109
|
+
class LlamaCppBackend:
|
|
110
|
+
"""The single model served by a llama-server instance."""
|
|
111
|
+
|
|
112
|
+
def __init__(self, spec: CandidateSpec, *, timeout: float = DEFAULT_TIMEOUT_SECONDS) -> None:
|
|
113
|
+
if spec.kind != "llamacpp":
|
|
114
|
+
raise SpecError(f"LlamaCppBackend cannot serve a {spec.kind!r} spec")
|
|
115
|
+
self._spec = spec
|
|
116
|
+
self._base_url = validate_base_url(spec.base_url)
|
|
117
|
+
self._timeout = timeout
|
|
118
|
+
self._info: ServerInfo | None = None
|
|
119
|
+
self._complete_token: int | None = None
|
|
120
|
+
|
|
121
|
+
@property
|
|
122
|
+
def spec(self) -> CandidateSpec:
|
|
123
|
+
return self._spec
|
|
124
|
+
|
|
125
|
+
def info(self) -> ServerInfo:
|
|
126
|
+
if self._info is None:
|
|
127
|
+
self._info = self._load_info()
|
|
128
|
+
return self._info
|
|
129
|
+
|
|
130
|
+
def chat(
|
|
131
|
+
self,
|
|
132
|
+
messages: Sequence[Message],
|
|
133
|
+
*,
|
|
134
|
+
max_tokens: int,
|
|
135
|
+
tools: Sequence[ToolSpec] = (),
|
|
136
|
+
json_schema: dict[str, JSONValue] | None = None,
|
|
137
|
+
seed: int = 0,
|
|
138
|
+
) -> ChatResult:
|
|
139
|
+
payload = openai_chat_payload(
|
|
140
|
+
self._spec.model,
|
|
141
|
+
messages,
|
|
142
|
+
max_tokens=max_tokens,
|
|
143
|
+
tools=tools,
|
|
144
|
+
json_schema=json_schema,
|
|
145
|
+
seed=seed,
|
|
146
|
+
)
|
|
147
|
+
started = time.perf_counter()
|
|
148
|
+
with explained_failures(self._explain):
|
|
149
|
+
body = post_json(self._url("/v1/chat/completions"), payload, timeout=self._timeout)
|
|
150
|
+
result = parse_openai_chat(body, seconds=time.perf_counter() - started)
|
|
151
|
+
return replace(result, decode_tokens_per_second=_predicted_rate(body))
|
|
152
|
+
|
|
153
|
+
def tokenize(self, text: str) -> tuple[int, ...]:
|
|
154
|
+
"""Tokenize with special tokens (such as BOS) added, exactly as a prompt would be."""
|
|
155
|
+
body = self._post("/tokenize", {"content": text, "add_special": True})
|
|
156
|
+
return tuple(_token_id(token) for token in expect_list(body.get("tokens"), "tokens"))
|
|
157
|
+
|
|
158
|
+
def generate_scored(self, prompt: str, *, max_tokens: int, top_k: int) -> list[TokenStep]:
|
|
159
|
+
"""Generate in one request, then token by token from the first folded entry on (see
|
|
160
|
+
the module docstring)."""
|
|
161
|
+
check_max_tokens(max_tokens)
|
|
162
|
+
check_top_k(top_k)
|
|
163
|
+
prompt_ids = self.tokenize(prompt)
|
|
164
|
+
batch = self._complete(prompt_ids, max_tokens=max_tokens, top_k=top_k)
|
|
165
|
+
steps = list(takewhile(lambda step: not folds_tokens(step), batch.steps))
|
|
166
|
+
if batch.stopped and len(steps) == len(batch.steps):
|
|
167
|
+
return steps
|
|
168
|
+
return self._extend(prompt_ids, steps, max_tokens=max_tokens, top_k=top_k)
|
|
169
|
+
|
|
170
|
+
def score_continuation(
|
|
171
|
+
self,
|
|
172
|
+
prompt: str,
|
|
173
|
+
continuation: Sequence[TokenStep],
|
|
174
|
+
*,
|
|
175
|
+
top_k: int,
|
|
176
|
+
prompt_token_ids: Sequence[int] | None = None,
|
|
177
|
+
) -> list[TopK]:
|
|
178
|
+
"""Feed token ids when the reference supplied them all, else fall back to text."""
|
|
179
|
+
check_top_k(top_k)
|
|
180
|
+
distributions: list[TopK] = []
|
|
181
|
+
for forced in _forced_prompts(prompt, continuation, prompt_token_ids):
|
|
182
|
+
steps = () if forced is None else self._next_token(forced, top_k=top_k).steps
|
|
183
|
+
distributions.append(steps[0].top if steps else ())
|
|
184
|
+
return distributions
|
|
185
|
+
|
|
186
|
+
def close(self) -> None:
|
|
187
|
+
"""Nothing to release: every request uses its own connection."""
|
|
188
|
+
|
|
189
|
+
def _extend(
|
|
190
|
+
self, prompt_ids: tuple[int, ...], steps: list[TokenStep], *, max_tokens: int, top_k: int
|
|
191
|
+
) -> list[TokenStep]:
|
|
192
|
+
"""Continue greedily one token per request until `max_tokens` steps or a stop.
|
|
193
|
+
|
|
194
|
+
A one-token reply never folds tokens, so every step is a single token with its own
|
|
195
|
+
distribution. Steps without ids (the older reply format) cannot be extended exactly
|
|
196
|
+
and are returned as they are.
|
|
197
|
+
"""
|
|
198
|
+
ids = list(prompt_ids)
|
|
199
|
+
for step in steps:
|
|
200
|
+
if step.chosen.token_id is None:
|
|
201
|
+
return steps
|
|
202
|
+
ids.append(step.chosen.token_id)
|
|
203
|
+
while len(steps) < max_tokens:
|
|
204
|
+
completion = self._next_token(tuple(ids), top_k=top_k)
|
|
205
|
+
if not completion.steps:
|
|
206
|
+
break
|
|
207
|
+
step = completion.steps[0]
|
|
208
|
+
steps.append(step)
|
|
209
|
+
if completion.stopped or step.chosen.token_id is None:
|
|
210
|
+
break
|
|
211
|
+
ids.append(step.chosen.token_id)
|
|
212
|
+
return steps
|
|
213
|
+
|
|
214
|
+
def _next_token(self, prompt: Prompt, *, top_k: int) -> _Completion:
|
|
215
|
+
"""Score the next token, recovering the distribution of a held-back token."""
|
|
216
|
+
completion = self._complete(prompt, max_tokens=1, top_k=top_k)
|
|
217
|
+
if completion.reported:
|
|
218
|
+
return completion
|
|
219
|
+
forced = self._complete(prompt, max_tokens=1, top_k=top_k, bias=self._complete_token_id())
|
|
220
|
+
top = forced.steps[0].top if forced.steps else ()
|
|
221
|
+
if not top:
|
|
222
|
+
raise CapabilityError(f"llama-server at {self._base_url} returned no probabilities")
|
|
223
|
+
return _Completion(steps=(TokenStep(chosen=top[0], top=top),), reported=True, stopped=False)
|
|
224
|
+
|
|
225
|
+
def _complete_token_id(self) -> int:
|
|
226
|
+
if self._complete_token is None:
|
|
227
|
+
body = self._post("/tokenize", {"content": _COMPLETE_TEXT, "add_special": False})
|
|
228
|
+
tokens = expect_list(body.get("tokens"), "tokens")
|
|
229
|
+
if not tokens:
|
|
230
|
+
raise BackendError(f"/tokenize returned no tokens for {_COMPLETE_TEXT!r}")
|
|
231
|
+
self._complete_token = _token_id(tokens[-1])
|
|
232
|
+
return self._complete_token
|
|
233
|
+
|
|
234
|
+
def _complete(
|
|
235
|
+
self, prompt: Prompt, *, max_tokens: int, top_k: int, bias: int | None = None
|
|
236
|
+
) -> _Completion:
|
|
237
|
+
payload: dict[str, JSONValue] = {
|
|
238
|
+
"prompt": prompt if isinstance(prompt, str) else list(prompt),
|
|
239
|
+
"n_predict": max_tokens,
|
|
240
|
+
"temperature": GREEDY_TEMPERATURE,
|
|
241
|
+
"n_probs": top_k,
|
|
242
|
+
"cache_prompt": True,
|
|
243
|
+
"seed": SCORING_SEED,
|
|
244
|
+
"samplers": list(_GREEDY_SAMPLERS),
|
|
245
|
+
}
|
|
246
|
+
if bias is not None:
|
|
247
|
+
payload["logit_bias"] = [[bias, _FORCING_BIAS]]
|
|
248
|
+
body = self._post("/completion", payload)
|
|
249
|
+
entries = body.get("completion_probabilities")
|
|
250
|
+
steps = (
|
|
251
|
+
()
|
|
252
|
+
if entries is None
|
|
253
|
+
else tuple(_parse_step(entry) for entry in expect_list(entries, "probabilities"))
|
|
254
|
+
)
|
|
255
|
+
return _Completion(steps=steps, reported=entries is not None, stopped=_stopped(body))
|
|
256
|
+
|
|
257
|
+
def _load_info(self) -> ServerInfo:
|
|
258
|
+
props = self._get("/props")
|
|
259
|
+
listing = self._get("/v1/models")
|
|
260
|
+
models = expect_list(listing.get("data") or [], "model list data")
|
|
261
|
+
listed = expect_dict(models[0], "model list entry") if models else {}
|
|
262
|
+
settings = props.get("default_generation_settings")
|
|
263
|
+
settings = settings if isinstance(settings, dict) else {}
|
|
264
|
+
model_file = _file_name(props.get("model_path"))
|
|
265
|
+
meta = listed.get("meta")
|
|
266
|
+
meta = meta if isinstance(meta, dict) else {}
|
|
267
|
+
facts = {
|
|
268
|
+
"model_file": model_file,
|
|
269
|
+
"trained_context_length": _optional_text(optional_int(meta.get("n_ctx_train"))),
|
|
270
|
+
"build": optional_str(props.get("build_info")),
|
|
271
|
+
}
|
|
272
|
+
# Size, parameter and vocabulary counts fingerprint the GGUF for the reference cache.
|
|
273
|
+
facts.update(
|
|
274
|
+
(f"model_{field}", _optional_text(optional_int(meta.get(field))))
|
|
275
|
+
for field in _META_FINGERPRINT
|
|
276
|
+
)
|
|
277
|
+
size, n_params, n_vocab = (optional_int(meta.get(field)) for field in _META_FINGERPRINT)
|
|
278
|
+
known = size is not None and n_params is not None and n_vocab is not None
|
|
279
|
+
return ServerInfo(
|
|
280
|
+
backend="llamacpp",
|
|
281
|
+
model=self._spec.model or optional_str(listed.get("id")) or model_file or "",
|
|
282
|
+
context_length=optional_int(settings.get("n_ctx")),
|
|
283
|
+
chat_template=optional_str(props.get("chat_template")),
|
|
284
|
+
template_dialect="jinja",
|
|
285
|
+
supports_logprobs=True,
|
|
286
|
+
exact_token_ids=True,
|
|
287
|
+
details=tuple((key, value) for key, value in facts.items() if value),
|
|
288
|
+
size_bytes=size,
|
|
289
|
+
weights_id=f"{size}:{n_params}:{n_vocab}" if known else None,
|
|
290
|
+
)
|
|
291
|
+
|
|
292
|
+
def _get(self, path: str) -> dict[str, JSONValue]:
|
|
293
|
+
with explained_failures(self._explain):
|
|
294
|
+
body = get_json(self._url(path), timeout=self._timeout)
|
|
295
|
+
return expect_dict(body, f"llama-server {path} response")
|
|
296
|
+
|
|
297
|
+
def _post(self, path: str, payload: dict[str, JSONValue]) -> dict[str, JSONValue]:
|
|
298
|
+
with explained_failures(self._explain):
|
|
299
|
+
body = post_json(self._url(path), payload, timeout=self._timeout)
|
|
300
|
+
return expect_dict(body, f"llama-server {path} response")
|
|
301
|
+
|
|
302
|
+
def _explain(self, failure: RequestFailure) -> str | None:
|
|
303
|
+
if failure.status is None:
|
|
304
|
+
port = urllib.parse.urlsplit(self._base_url).port or _DEFAULT_PORT
|
|
305
|
+
return (
|
|
306
|
+
f"cannot reach llama-server at {self._base_url} ({failure.detail}); "
|
|
307
|
+
f"start it with `llama-server -m model.gguf --port {port}`"
|
|
308
|
+
)
|
|
309
|
+
if failure.status == _LOADING_STATUS and "loading" in failure.detail.lower():
|
|
310
|
+
return (
|
|
311
|
+
f"llama-server at {self._base_url} is still loading the model; "
|
|
312
|
+
"wait until its log says the server is listening, then run again"
|
|
313
|
+
)
|
|
314
|
+
if failure.status == 404:
|
|
315
|
+
return (
|
|
316
|
+
f"{printable(failure.url)} returned HTTP 404; "
|
|
317
|
+
f"check that {self._base_url} is a llama-server and not another kind of server"
|
|
318
|
+
)
|
|
319
|
+
return None
|
|
320
|
+
|
|
321
|
+
def _url(self, path: str) -> str:
|
|
322
|
+
return self._base_url + path
|
|
323
|
+
|
|
324
|
+
|
|
325
|
+
def _forced_prompts(
|
|
326
|
+
prompt: str, continuation: Sequence[TokenStep], prompt_token_ids: Sequence[int] | None
|
|
327
|
+
) -> list[Prompt | None]:
|
|
328
|
+
"""One prompt per continuation position: token ids when every id is known, else text.
|
|
329
|
+
|
|
330
|
+
None marks a position that text cannot reproduce; see `forced_texts`.
|
|
331
|
+
"""
|
|
332
|
+
chosen_ids = [step.chosen.token_id for step in continuation]
|
|
333
|
+
known = tuple(token_id for token_id in chosen_ids if token_id is not None)
|
|
334
|
+
if prompt_token_ids is None or len(known) != len(chosen_ids):
|
|
335
|
+
return list(forced_texts(prompt, continuation))
|
|
336
|
+
prompt_ids = tuple(prompt_token_ids)
|
|
337
|
+
return [prompt_ids + known[:position] for position in range(len(known))]
|
|
338
|
+
|
|
339
|
+
|
|
340
|
+
def _stopped(body: dict[str, JSONValue]) -> bool:
|
|
341
|
+
"""Read `stop_type`, or the `stopped_eos` and `stopped_word` flags of older servers."""
|
|
342
|
+
return (
|
|
343
|
+
body.get("stop_type") in _STOP_TYPES
|
|
344
|
+
or body.get("stopped_eos") is True
|
|
345
|
+
or body.get("stopped_word") is True
|
|
346
|
+
)
|
|
347
|
+
|
|
348
|
+
|
|
349
|
+
def _predicted_rate(body: JSONValue) -> float | None:
|
|
350
|
+
"""Read the decode timing llama-server adds to chat responses, when present.
|
|
351
|
+
|
|
352
|
+
Computed from predicted_n and predicted_ms (which is how the server derives its own
|
|
353
|
+
predicted_per_second) so a one-token answer can be recognized and left out.
|
|
354
|
+
"""
|
|
355
|
+
timings = body.get("timings") if isinstance(body, dict) else None
|
|
356
|
+
if not isinstance(timings, dict):
|
|
357
|
+
return None
|
|
358
|
+
milliseconds = timings.get("predicted_ms")
|
|
359
|
+
if isinstance(milliseconds, bool) or not isinstance(milliseconds, (int, float)):
|
|
360
|
+
return None
|
|
361
|
+
return decode_rate(optional_int(timings.get("predicted_n")), milliseconds / 1000)
|
|
362
|
+
|
|
363
|
+
|
|
364
|
+
def _optional_text(value: int | None) -> str | None:
|
|
365
|
+
return None if value is None else str(value)
|
|
366
|
+
|
|
367
|
+
|
|
368
|
+
def _token_id(token: JSONValue) -> int:
|
|
369
|
+
"""Read one /tokenize entry: a bare id, or {"id", "piece"} when pieces are requested."""
|
|
370
|
+
value = token.get("id") if isinstance(token, dict) else token
|
|
371
|
+
token_id = optional_int(value)
|
|
372
|
+
if token_id is None:
|
|
373
|
+
raise BackendError(f"/tokenize returned an invalid token entry: {token!r}")
|
|
374
|
+
return token_id
|
|
375
|
+
|
|
376
|
+
|
|
377
|
+
def _file_name(path: JSONValue) -> str | None:
|
|
378
|
+
if not isinstance(path, str):
|
|
379
|
+
return None
|
|
380
|
+
# The server may run on Windows or POSIX; PureWindowsPath splits on both separators.
|
|
381
|
+
return PureWindowsPath(path).name or None
|
|
382
|
+
|
|
383
|
+
|
|
384
|
+
def _parse_step(entry: JSONValue) -> TokenStep:
|
|
385
|
+
step = expect_dict(entry, "completion_probabilities entry")
|
|
386
|
+
if "probs" in step:
|
|
387
|
+
return _parse_legacy_step(step)
|
|
388
|
+
top = expect_list(step.get("top_logprobs") or [], "top_logprobs")
|
|
389
|
+
alternatives = (_parse_prob(item) for item in top if not _is_impossible(item))
|
|
390
|
+
return TokenStep(chosen=_parse_prob(step), top=sorted_top(alternatives))
|
|
391
|
+
|
|
392
|
+
|
|
393
|
+
def _is_impossible(item: JSONValue) -> bool:
|
|
394
|
+
"""True for an alternative whose logprob is null, the server's encoding of -inf."""
|
|
395
|
+
return isinstance(item, dict) and "logprob" in item and item["logprob"] is None
|
|
396
|
+
|
|
397
|
+
|
|
398
|
+
def _parse_prob(value: JSONValue) -> TokenProb:
|
|
399
|
+
item = expect_dict(value, "token probability")
|
|
400
|
+
return TokenProb(
|
|
401
|
+
token=expect_str(item.get("token"), "token"),
|
|
402
|
+
logprob=expect_float(item.get("logprob"), "logprob"),
|
|
403
|
+
token_id=optional_int(item.get("id")),
|
|
404
|
+
token_bytes=optional_bytes(item.get("bytes")),
|
|
405
|
+
)
|
|
406
|
+
|
|
407
|
+
|
|
408
|
+
def _parse_legacy_step(step: dict[str, JSONValue]) -> TokenStep:
|
|
409
|
+
"""Read the older format, which reports probabilities and no token ids."""
|
|
410
|
+
content = expect_str(step.get("content"), "content")
|
|
411
|
+
probs = (_parse_legacy_prob(item) for item in expect_list(step["probs"], "probs"))
|
|
412
|
+
top = sorted_top(prob for prob in probs if prob is not None)
|
|
413
|
+
chosen = next((prob for prob in top if prob.token == content), None)
|
|
414
|
+
if chosen is None:
|
|
415
|
+
raise BackendError(f"chosen token {content!r} is missing from its own probabilities")
|
|
416
|
+
return TokenStep(chosen=chosen, top=top)
|
|
417
|
+
|
|
418
|
+
|
|
419
|
+
def _parse_legacy_prob(value: JSONValue) -> TokenProb | None:
|
|
420
|
+
"""Return None for zero-probability filler entries, which have no logarithm."""
|
|
421
|
+
item = expect_dict(value, "probs entry")
|
|
422
|
+
prob = expect_float(item.get("prob"), "prob")
|
|
423
|
+
if prob == 0.0:
|
|
424
|
+
return None
|
|
425
|
+
return TokenProb(
|
|
426
|
+
token=expect_str(item.get("tok_str"), "tok_str"),
|
|
427
|
+
logprob=logprob_from_prob(prob, "prob"),
|
|
428
|
+
)
|