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,342 @@
|
|
|
1
|
+
"""Helpers shared by the backend adapters: payload validation, OpenAI chat format, and
|
|
2
|
+
turning raw HTTP failures into messages that tell the user what to do."""
|
|
3
|
+
|
|
4
|
+
from __future__ import annotations
|
|
5
|
+
|
|
6
|
+
import json
|
|
7
|
+
import math
|
|
8
|
+
from collections.abc import Callable, Iterable, Iterator, Sequence
|
|
9
|
+
from contextlib import contextmanager
|
|
10
|
+
from dataclasses import dataclass
|
|
11
|
+
from typing import Final
|
|
12
|
+
|
|
13
|
+
from quantdiff._text import printable
|
|
14
|
+
from quantdiff.backends.base import MAX_TOP_K
|
|
15
|
+
from quantdiff.errors import BackendError, CapabilityError, RequestError, SpecError
|
|
16
|
+
from quantdiff.types import (
|
|
17
|
+
ChatResult,
|
|
18
|
+
JSONValue,
|
|
19
|
+
Message,
|
|
20
|
+
TokenProb,
|
|
21
|
+
TokenStep,
|
|
22
|
+
ToolCall,
|
|
23
|
+
ToolSpec,
|
|
24
|
+
TopK,
|
|
25
|
+
)
|
|
26
|
+
|
|
27
|
+
GREEDY_TEMPERATURE: Final = 0.0
|
|
28
|
+
SCORING_SEED: Final = 0
|
|
29
|
+
"""Seed for raw scoring requests. Greedy decoding ignores it, but some servers want one."""
|
|
30
|
+
_MAX_BYTE: Final = 255
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
# Request arguments -----------------------------------------------------------------------
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def check_top_k(top_k: int) -> None:
|
|
37
|
+
if not 1 <= top_k <= MAX_TOP_K:
|
|
38
|
+
raise SpecError(f"top_k must be between 1 and {MAX_TOP_K}, got {top_k}")
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def check_max_tokens(max_tokens: int) -> None:
|
|
42
|
+
if max_tokens < 1:
|
|
43
|
+
raise SpecError(f"max_tokens must be at least 1, got {max_tokens}")
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
def forced_texts(prompt: str, continuation: Sequence[TokenStep]) -> list[str | None]:
|
|
47
|
+
"""Return the prompt text that teacher-forces each continuation position, or None for a
|
|
48
|
+
position that text cannot reproduce.
|
|
49
|
+
|
|
50
|
+
Prefixes are rebuilt from each chosen token's bytes when the server reported them,
|
|
51
|
+
because a token that holds only part of a UTF-8 character has no exact text of its own;
|
|
52
|
+
tokens without bytes contribute their text. A position is None when the reference left
|
|
53
|
+
it unscored (its `top` is empty) or when its prefix ends inside a character, since a
|
|
54
|
+
prompt must be valid text.
|
|
55
|
+
"""
|
|
56
|
+
texts: list[str | None] = []
|
|
57
|
+
prefix = bytearray(prompt.encode("utf-8"))
|
|
58
|
+
for step in continuation:
|
|
59
|
+
texts.append(_decoded(prefix) if step.top else None)
|
|
60
|
+
chosen = step.chosen
|
|
61
|
+
prefix += chosen.token.encode("utf-8") if chosen.token_bytes is None else chosen.token_bytes
|
|
62
|
+
return texts
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
def _decoded(data: bytearray) -> str | None:
|
|
66
|
+
try:
|
|
67
|
+
return data.decode("utf-8")
|
|
68
|
+
except UnicodeDecodeError:
|
|
69
|
+
return None
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
# Folded tokens ---------------------------------------------------------------------------
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def unfolded_steps(steps: Iterable[TokenStep]) -> list[TokenStep]:
|
|
76
|
+
"""Keep every step, but empty the `top` of each step that folds several tokens.
|
|
77
|
+
|
|
78
|
+
Ollama and llama-server hold back a token that ends inside a UTF-8 character and report
|
|
79
|
+
it together with the tokens that complete the character, as one entry whose text is
|
|
80
|
+
the whole character but whose logprob and alternatives belong to the last of those
|
|
81
|
+
tokens. That distribution sits one or more tokens after the position the entry stands
|
|
82
|
+
for, so comparing it with a teacher-forced candidate would compare different
|
|
83
|
+
positions. An empty `top` marks the position as unscored; metrics skip it.
|
|
84
|
+
"""
|
|
85
|
+
return [TokenStep(step.chosen, ()) if folds_tokens(step) else step for step in steps]
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
def folds_tokens(step: TokenStep) -> bool:
|
|
89
|
+
"""True when the server folded several tokens into this step (see `unfolded_steps`).
|
|
90
|
+
|
|
91
|
+
Greedy decoding picks the head of the distribution, so a step's own token is always in
|
|
92
|
+
its `top`. A folded step's text is missing from that list, which holds only the
|
|
93
|
+
character fragments the last folded token chose between, or is listed with different
|
|
94
|
+
bytes when the server matches it by id.
|
|
95
|
+
"""
|
|
96
|
+
chosen = step.chosen
|
|
97
|
+
own = next((prob for prob in step.top if _same_token(prob, chosen)), None)
|
|
98
|
+
if own is None:
|
|
99
|
+
return bool(step.top)
|
|
100
|
+
if own.token_bytes is None or chosen.token_bytes is None:
|
|
101
|
+
return False
|
|
102
|
+
return own.token_bytes != chosen.token_bytes
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
def _same_token(prob: TokenProb, chosen: TokenProb) -> bool:
|
|
106
|
+
if prob.token_id is not None and chosen.token_id is not None:
|
|
107
|
+
return prob.token_id == chosen.token_id
|
|
108
|
+
return prob.token == chosen.token
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
# Response validation ---------------------------------------------------------------------
|
|
112
|
+
|
|
113
|
+
|
|
114
|
+
def expect_dict(value: JSONValue, where: str) -> dict[str, JSONValue]:
|
|
115
|
+
if not isinstance(value, dict):
|
|
116
|
+
raise BackendError(f"{where}: expected a JSON object, got {type(value).__name__}")
|
|
117
|
+
return value
|
|
118
|
+
|
|
119
|
+
|
|
120
|
+
def expect_list(value: JSONValue, where: str) -> list[JSONValue]:
|
|
121
|
+
if not isinstance(value, list):
|
|
122
|
+
raise BackendError(f"{where}: expected a JSON array, got {type(value).__name__}")
|
|
123
|
+
return value
|
|
124
|
+
|
|
125
|
+
|
|
126
|
+
def expect_str(value: JSONValue, where: str) -> str:
|
|
127
|
+
if not isinstance(value, str):
|
|
128
|
+
raise BackendError(f"{where}: expected a string, got {type(value).__name__}")
|
|
129
|
+
return value
|
|
130
|
+
|
|
131
|
+
|
|
132
|
+
def expect_float(value: JSONValue, where: str) -> float:
|
|
133
|
+
"""Return a finite number. Metrics cannot use NaN or infinity, and json.loads turns an
|
|
134
|
+
out-of-range literal such as 1e999 into infinity."""
|
|
135
|
+
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
|
136
|
+
raise BackendError(f"{where}: expected a number, got {type(value).__name__}")
|
|
137
|
+
number = float(value)
|
|
138
|
+
if not math.isfinite(number):
|
|
139
|
+
raise BackendError(f"{where}: expected a finite number, got {number}")
|
|
140
|
+
return number
|
|
141
|
+
|
|
142
|
+
|
|
143
|
+
def optional_int(value: JSONValue) -> int | None:
|
|
144
|
+
"""Return `value` if it is a JSON integer, else None. Booleans are not integers."""
|
|
145
|
+
if isinstance(value, int) and not isinstance(value, bool):
|
|
146
|
+
return value
|
|
147
|
+
return None
|
|
148
|
+
|
|
149
|
+
|
|
150
|
+
def optional_str(value: JSONValue) -> str | None:
|
|
151
|
+
return value if isinstance(value, str) else None
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
def optional_bytes(value: JSONValue) -> bytes | None:
|
|
155
|
+
"""Read a token's `bytes` field, a JSON array of byte values, or return None if absent
|
|
156
|
+
or malformed."""
|
|
157
|
+
if not isinstance(value, list):
|
|
158
|
+
return None
|
|
159
|
+
if not all(isinstance(item, int) and not isinstance(item, bool) for item in value):
|
|
160
|
+
return None
|
|
161
|
+
if not all(0 <= item <= _MAX_BYTE for item in value):
|
|
162
|
+
return None
|
|
163
|
+
return bytes(value)
|
|
164
|
+
|
|
165
|
+
|
|
166
|
+
def decode_rate(tokens: int | None, seconds: float | None) -> float | None:
|
|
167
|
+
"""Return the server-measured decode speed, or None when it is not meaningful.
|
|
168
|
+
|
|
169
|
+
A lone generated token is sampled from the prompt-processing pass, so servers report
|
|
170
|
+
a decode time near zero for it and the quotient would be absurd.
|
|
171
|
+
"""
|
|
172
|
+
if tokens is None or tokens < 2 or seconds is None or seconds <= 0:
|
|
173
|
+
return None
|
|
174
|
+
rate = tokens / seconds
|
|
175
|
+
return rate if math.isfinite(rate) else None
|
|
176
|
+
|
|
177
|
+
|
|
178
|
+
def sorted_top(probs: Iterable[TokenProb]) -> TopK:
|
|
179
|
+
"""Sort by descending logprob, keeping server order for ties."""
|
|
180
|
+
return tuple(sorted(probs, key=lambda prob: prob.logprob, reverse=True))
|
|
181
|
+
|
|
182
|
+
|
|
183
|
+
def logprob_from_prob(prob: float, where: str) -> float:
|
|
184
|
+
if not 0.0 < prob <= 1.0:
|
|
185
|
+
raise BackendError(f"{where}: probability {prob} is outside (0, 1]")
|
|
186
|
+
return math.log(prob)
|
|
187
|
+
|
|
188
|
+
|
|
189
|
+
# OpenAI chat format ----------------------------------------------------------------------
|
|
190
|
+
|
|
191
|
+
|
|
192
|
+
def openai_tool(tool: ToolSpec) -> dict[str, JSONValue]:
|
|
193
|
+
return {
|
|
194
|
+
"type": "function",
|
|
195
|
+
"function": {
|
|
196
|
+
"name": tool.name,
|
|
197
|
+
"description": tool.description,
|
|
198
|
+
"parameters": tool.parameters,
|
|
199
|
+
},
|
|
200
|
+
}
|
|
201
|
+
|
|
202
|
+
|
|
203
|
+
def openai_messages(messages: Sequence[Message]) -> list[JSONValue]:
|
|
204
|
+
return [{"role": message.role, "content": message.content} for message in messages]
|
|
205
|
+
|
|
206
|
+
|
|
207
|
+
def openai_chat_payload(
|
|
208
|
+
model: str,
|
|
209
|
+
messages: Sequence[Message],
|
|
210
|
+
*,
|
|
211
|
+
max_tokens: int,
|
|
212
|
+
tools: Sequence[ToolSpec],
|
|
213
|
+
json_schema: dict[str, JSONValue] | None,
|
|
214
|
+
seed: int,
|
|
215
|
+
) -> dict[str, JSONValue]:
|
|
216
|
+
"""Build a greedy /chat/completions request. An empty `model` is left out."""
|
|
217
|
+
check_max_tokens(max_tokens)
|
|
218
|
+
payload: dict[str, JSONValue] = {
|
|
219
|
+
"messages": openai_messages(messages),
|
|
220
|
+
"max_tokens": max_tokens,
|
|
221
|
+
"temperature": GREEDY_TEMPERATURE,
|
|
222
|
+
"seed": seed,
|
|
223
|
+
"stream": False,
|
|
224
|
+
}
|
|
225
|
+
if model:
|
|
226
|
+
payload["model"] = model
|
|
227
|
+
if tools:
|
|
228
|
+
payload["tools"] = [openai_tool(tool) for tool in tools]
|
|
229
|
+
if json_schema is not None:
|
|
230
|
+
payload["response_format"] = {
|
|
231
|
+
"type": "json_schema",
|
|
232
|
+
"json_schema": {"name": "answer", "schema": json_schema},
|
|
233
|
+
}
|
|
234
|
+
return payload
|
|
235
|
+
|
|
236
|
+
|
|
237
|
+
def parse_openai_chat(body: JSONValue, *, seconds: float) -> ChatResult:
|
|
238
|
+
"""Read the first choice of a /chat/completions response."""
|
|
239
|
+
response = expect_dict(body, "chat response")
|
|
240
|
+
choices = expect_list(response.get("choices"), "chat response choices")
|
|
241
|
+
if not choices:
|
|
242
|
+
raise BackendError("chat response has no choices")
|
|
243
|
+
choice = expect_dict(choices[0], "chat choice")
|
|
244
|
+
message = expect_dict(choice.get("message"), "chat message")
|
|
245
|
+
raw_calls = message.get("tool_calls") or []
|
|
246
|
+
usage = response.get("usage")
|
|
247
|
+
usage = usage if isinstance(usage, dict) else {}
|
|
248
|
+
return ChatResult(
|
|
249
|
+
text=optional_str(message.get("content")) or "",
|
|
250
|
+
tool_calls=tuple(
|
|
251
|
+
parse_openai_tool_call(call) for call in expect_list(raw_calls, "tool_calls")
|
|
252
|
+
),
|
|
253
|
+
finish_reason=optional_str(choice.get("finish_reason")),
|
|
254
|
+
prompt_tokens=optional_int(usage.get("prompt_tokens")),
|
|
255
|
+
completion_tokens=optional_int(usage.get("completion_tokens")),
|
|
256
|
+
seconds=seconds,
|
|
257
|
+
)
|
|
258
|
+
|
|
259
|
+
|
|
260
|
+
def parse_openai_tool_call(call: JSONValue) -> ToolCall:
|
|
261
|
+
function = expect_dict(expect_dict(call, "tool call").get("function"), "tool call function")
|
|
262
|
+
name = expect_str(function.get("name"), "tool call name")
|
|
263
|
+
return tool_call(name, function.get("arguments"))
|
|
264
|
+
|
|
265
|
+
|
|
266
|
+
def tool_call(name: str, arguments: JSONValue) -> ToolCall:
|
|
267
|
+
"""Build a ToolCall from arguments sent either as a JSON string or as an object."""
|
|
268
|
+
if isinstance(arguments, str):
|
|
269
|
+
return ToolCall(name=name, arguments=parse_arguments(arguments), raw_arguments=arguments)
|
|
270
|
+
raw = json.dumps(arguments)
|
|
271
|
+
return ToolCall(
|
|
272
|
+
name=name,
|
|
273
|
+
arguments=arguments if isinstance(arguments, dict) else None,
|
|
274
|
+
raw_arguments=raw,
|
|
275
|
+
)
|
|
276
|
+
|
|
277
|
+
|
|
278
|
+
def parse_arguments(raw: str) -> dict[str, JSONValue] | None:
|
|
279
|
+
"""Decode tool arguments, or return None when they are not a JSON object."""
|
|
280
|
+
try:
|
|
281
|
+
decoded = json.loads(raw)
|
|
282
|
+
except json.JSONDecodeError:
|
|
283
|
+
return None
|
|
284
|
+
return decoded if isinstance(decoded, dict) else None
|
|
285
|
+
|
|
286
|
+
|
|
287
|
+
# Friendly failures -----------------------------------------------------------------------
|
|
288
|
+
|
|
289
|
+
|
|
290
|
+
@dataclass(frozen=True, slots=True)
|
|
291
|
+
class RequestFailure:
|
|
292
|
+
"""A failed request, taken from the RequestError that quantdiff._http raised."""
|
|
293
|
+
|
|
294
|
+
url: str
|
|
295
|
+
status: int | None
|
|
296
|
+
"""HTTP status, or None when the server could not be reached at all."""
|
|
297
|
+
detail: str
|
|
298
|
+
"""Response body preview for status errors, else the transport error reason."""
|
|
299
|
+
|
|
300
|
+
|
|
301
|
+
def request_failure(exc: BackendError) -> RequestFailure | None:
|
|
302
|
+
"""Classify `exc`, or return None for failures that are not about reaching the server."""
|
|
303
|
+
if not isinstance(exc, RequestError):
|
|
304
|
+
return None
|
|
305
|
+
if exc.status is not None:
|
|
306
|
+
return RequestFailure(exc.url, exc.status, exc.detail)
|
|
307
|
+
return RequestFailure(exc.url, None, _short_reason(exc.detail))
|
|
308
|
+
|
|
309
|
+
|
|
310
|
+
Explainer = Callable[[RequestFailure], str | None]
|
|
311
|
+
|
|
312
|
+
|
|
313
|
+
@contextmanager
|
|
314
|
+
def explained_failures(explain: Explainer) -> Iterator[None]:
|
|
315
|
+
"""Re-raise request failures that `explain` recognizes as a BackendError with its message.
|
|
316
|
+
|
|
317
|
+
The original error stays available as `__cause__`. Failures `explain` does not recognize
|
|
318
|
+
propagate unchanged.
|
|
319
|
+
"""
|
|
320
|
+
try:
|
|
321
|
+
yield
|
|
322
|
+
except CapabilityError:
|
|
323
|
+
raise
|
|
324
|
+
except BackendError as exc:
|
|
325
|
+
failure = request_failure(exc)
|
|
326
|
+
message = None if failure is None else explain(failure)
|
|
327
|
+
if message is None:
|
|
328
|
+
raise
|
|
329
|
+
raise BackendError(message) from exc
|
|
330
|
+
|
|
331
|
+
|
|
332
|
+
def mentions_missing_model(detail: str) -> bool:
|
|
333
|
+
"""True for error bodies such as `model 'x' not found` or `The model x does not exist`."""
|
|
334
|
+
lowered = detail.lower()
|
|
335
|
+
return "model" in lowered and ("not found" in lowered or "does not exist" in lowered)
|
|
336
|
+
|
|
337
|
+
|
|
338
|
+
def _short_reason(reason: str) -> str:
|
|
339
|
+
# Windows reports a refused connection as a long WinError sentence.
|
|
340
|
+
if "refused" in reason.lower():
|
|
341
|
+
return "connection refused"
|
|
342
|
+
return printable(reason)
|
|
@@ -0,0 +1,91 @@
|
|
|
1
|
+
"""The contract every model server adapter implements."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from collections.abc import Sequence
|
|
6
|
+
from typing import Protocol, runtime_checkable
|
|
7
|
+
|
|
8
|
+
from quantdiff.types import (
|
|
9
|
+
CandidateSpec,
|
|
10
|
+
ChatResult,
|
|
11
|
+
JSONValue,
|
|
12
|
+
Message,
|
|
13
|
+
ServerInfo,
|
|
14
|
+
TokenStep,
|
|
15
|
+
ToolSpec,
|
|
16
|
+
TopK,
|
|
17
|
+
)
|
|
18
|
+
|
|
19
|
+
MAX_TOP_K = 20
|
|
20
|
+
"""Upper bound on top-k logprobs. Ollama and OpenAI-compatible servers cap at 20."""
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
@runtime_checkable
|
|
24
|
+
class Backend(Protocol):
|
|
25
|
+
"""A connection to one served model.
|
|
26
|
+
|
|
27
|
+
All generation is greedy (temperature 0) so runs are as repeatable as the server
|
|
28
|
+
allows. Implementations raise BackendError on transport or protocol failures and
|
|
29
|
+
CapabilityError when an operation is unsupported, never bare exceptions.
|
|
30
|
+
"""
|
|
31
|
+
|
|
32
|
+
@property
|
|
33
|
+
def spec(self) -> CandidateSpec: ...
|
|
34
|
+
|
|
35
|
+
def info(self) -> ServerInfo:
|
|
36
|
+
"""Describe the served model. May be called many times; implementations cache it."""
|
|
37
|
+
...
|
|
38
|
+
|
|
39
|
+
def chat(
|
|
40
|
+
self,
|
|
41
|
+
messages: Sequence[Message],
|
|
42
|
+
*,
|
|
43
|
+
max_tokens: int,
|
|
44
|
+
tools: Sequence[ToolSpec] = (),
|
|
45
|
+
json_schema: dict[str, JSONValue] | None = None,
|
|
46
|
+
seed: int = 0,
|
|
47
|
+
) -> ChatResult:
|
|
48
|
+
"""Run one chat completion through the server's own chat template."""
|
|
49
|
+
...
|
|
50
|
+
|
|
51
|
+
def tokenize(self, text: str) -> tuple[int, ...] | None:
|
|
52
|
+
"""Return token ids for raw `text`, or None if the backend cannot tokenize."""
|
|
53
|
+
...
|
|
54
|
+
|
|
55
|
+
def generate_scored(self, prompt: str, *, max_tokens: int, top_k: int) -> list[TokenStep]:
|
|
56
|
+
"""Greedily continue raw `prompt` (no chat template), returning top-k at each step.
|
|
57
|
+
|
|
58
|
+
A step whose `top` is empty marks a position no candidate can be compared at, such
|
|
59
|
+
as one whose reported distribution belongs to a later token; metrics skip it.
|
|
60
|
+
`chosen` still holds the token, so teacher forcing can continue past it.
|
|
61
|
+
|
|
62
|
+
Raises CapabilityError if the backend does not expose logprobs.
|
|
63
|
+
"""
|
|
64
|
+
...
|
|
65
|
+
|
|
66
|
+
def score_continuation(
|
|
67
|
+
self,
|
|
68
|
+
prompt: str,
|
|
69
|
+
continuation: Sequence[TokenStep],
|
|
70
|
+
*,
|
|
71
|
+
top_k: int,
|
|
72
|
+
prompt_token_ids: Sequence[int] | None = None,
|
|
73
|
+
) -> list[TopK]:
|
|
74
|
+
"""Teacher-force `continuation` after raw `prompt` and return the top-k next-token
|
|
75
|
+
distribution at every position, so result[i] is conditioned on continuation[:i].
|
|
76
|
+
|
|
77
|
+
When `info().exact_token_ids` is True the implementation must feed token ids
|
|
78
|
+
(`prompt_token_ids` plus each step's `chosen.token_id`) instead of text, so the
|
|
79
|
+
candidate sees exactly the reference's token sequence.
|
|
80
|
+
|
|
81
|
+
result[i] is empty when the server stopped there without a distribution, or when
|
|
82
|
+
the position cannot be scored: the reference left it unscored (empty `top`, sent
|
|
83
|
+
without a request), or a text prompt cannot reproduce it.
|
|
84
|
+
|
|
85
|
+
Raises CapabilityError if the backend does not expose logprobs.
|
|
86
|
+
"""
|
|
87
|
+
...
|
|
88
|
+
|
|
89
|
+
def close(self) -> None:
|
|
90
|
+
"""Release resources. Safe to call more than once."""
|
|
91
|
+
...
|