logit-classifier 0.1.0__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.
@@ -0,0 +1,90 @@
1
+ """Local zero-shot classifier that reads its answer off a logit row.
2
+
3
+ One state and a set of declared questions go in. One calibrated probability per
4
+ declared option comes out. No token is generated.
5
+
6
+ Importing this package pulls in numpy alone. The transformers backend arrives
7
+ with the `[hf]` extra and the HTTP service with `[service]`.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ from .backends.base import (
13
+ Backend,
14
+ BackendContractError,
15
+ BranchLogits,
16
+ VisionUnsupportedError,
17
+ verify_backend,
18
+ )
19
+ from .classifier import Classifier, Diagnostics, load_model
20
+ from .config import ANSWER_PREFILL, Config
21
+ from .deps import MissingDependencyError
22
+ from .errors import ConfigError, LogitClassifierError
23
+ from .labels import (
24
+ LABEL_ALPHABET,
25
+ MAX_LABELS_PER_BRANCH,
26
+ LabelBoundaryError,
27
+ verify_label_ids,
28
+ )
29
+ from .schema import (
30
+ Answer,
31
+ ChoiceAnswer,
32
+ ChoiceQuestion,
33
+ NoulAnswer,
34
+ NoulCriteria,
35
+ NoulQuestion,
36
+ Question,
37
+ SchemaError,
38
+ ScoreAnswer,
39
+ ScoreQuestion,
40
+ SystemOneRequest,
41
+ SystemOneResponse,
42
+ Usage,
43
+ parse_questions,
44
+ parse_request,
45
+ )
46
+ from .vision import ImageError
47
+
48
+ __version__ = "0.1.0"
49
+
50
+ # The ComfyUI socket type a node pack declares for a loaded classifier. It lives
51
+ # here so the library and the packs cannot drift apart on the spelling.
52
+ COMFY_SOCKET_TYPE = "LOGIT_CLASSIFIER"
53
+
54
+ __all__ = [
55
+ "ANSWER_PREFILL",
56
+ "COMFY_SOCKET_TYPE",
57
+ "LABEL_ALPHABET",
58
+ "MAX_LABELS_PER_BRANCH",
59
+ "Answer",
60
+ "Backend",
61
+ "BackendContractError",
62
+ "BranchLogits",
63
+ "ChoiceAnswer",
64
+ "ChoiceQuestion",
65
+ "Classifier",
66
+ "Config",
67
+ "ConfigError",
68
+ "Diagnostics",
69
+ "ImageError",
70
+ "LabelBoundaryError",
71
+ "LogitClassifierError",
72
+ "MissingDependencyError",
73
+ "NoulAnswer",
74
+ "NoulCriteria",
75
+ "NoulQuestion",
76
+ "Question",
77
+ "SchemaError",
78
+ "ScoreAnswer",
79
+ "ScoreQuestion",
80
+ "SystemOneRequest",
81
+ "SystemOneResponse",
82
+ "Usage",
83
+ "VisionUnsupportedError",
84
+ "__version__",
85
+ "load_model",
86
+ "parse_questions",
87
+ "parse_request",
88
+ "verify_backend",
89
+ "verify_label_ids",
90
+ ]
@@ -0,0 +1,10 @@
1
+ """`python -m logit_classifier`, for a venv whose scripts directory is not on PATH."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import sys
6
+
7
+ from .cli import main
8
+
9
+ if __name__ == "__main__":
10
+ sys.exit(main(prog="python -m logit_classifier"))
@@ -0,0 +1,5 @@
1
+ """Readout backends, one per host that can run a forward pass.
2
+
3
+ Importing this package pulls in no host. `hf` needs the `[hf]` extra, so it is
4
+ imported by name at the point of use. `base` is the single path to the port.
5
+ """
@@ -0,0 +1,125 @@
1
+ """The readout port, the one boundary between the arithmetic and a host's forward pass.
2
+
3
+ Everything above this port is stdlib and numpy. Everything below it is one host's
4
+ way of turning token ids into a logit row: transformers here, a ComfyUI CLIP next.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ from dataclasses import dataclass
10
+ from typing import Any, Protocol, runtime_checkable
11
+
12
+ import numpy as np
13
+
14
+ from ..config import ANSWER_PREFILL
15
+ from ..errors import LogitClassifierError
16
+ from ..labels import MAX_LABELS_PER_BRANCH, verify_label_ids
17
+
18
+ # The probe text only has to be stable. Routing the label proof through render
19
+ # rather than a hand-built template measured identical ids on both shipped models.
20
+ _PROBE_SYSTEM = "probe system"
21
+ _PROBE_STATE = "probe state"
22
+ # The closed render must carry a branch suffix, since _suffix_ids splits the real
23
+ # closed render on the open-ended one and a shared body would not exercise that.
24
+ _PROBE_SUFFIX = "\nQuestion: (A) probe option"
25
+
26
+
27
+ class VisionUnsupportedError(LogitClassifierError, ValueError):
28
+ """An image was given to a model that carries no vision tower."""
29
+
30
+
31
+ class BackendContractError(LogitClassifierError, RuntimeError):
32
+ """A backend returned rows the port does not allow."""
33
+
34
+
35
+ @dataclass(frozen=True)
36
+ class BranchLogits:
37
+ """Raw label logits for one branch, before any calibration."""
38
+
39
+ #: One logit per label this branch asked for, ordered as LABEL_ALPHABET[:count]
40
+ #: and exactly that wide, never the full vocabulary row.
41
+ z: np.ndarray
42
+ # Share of full-vocabulary probability mass sitting on the option labels. A
43
+ # low value means the restricted softmax is normalising noise.
44
+ candidate_mass: float
45
+
46
+
47
+ @runtime_checkable
48
+ class Backend(Protocol):
49
+ """What a host must provide for the classifier to read a distribution off it.
50
+
51
+ Runtime checkable, so a node pack can assert its own duck-typed object
52
+ satisfies the port before it reaches the classifier.
53
+
54
+ A backend may also declare `canonical_model_id: str`, the identity the fitted
55
+ temperature and the calibration fingerprint key on. Without it, `model_id` is that
56
+ identity.
57
+ """
58
+
59
+ #: The key FITTED_TEMPERATURES is looked up by when no canonical_model_id is declared.
60
+ #: Crossing the two shipped models took calibration error from 0.088 to 0.565, so a
61
+ #: wrong string is silently miscalibrated.
62
+ model_id: str
63
+ #: One token id per label in LABEL_ALPHABET, proven at load against the rendered prefill.
64
+ label_ids: list[int]
65
+ #: Whether this host's model carries a vision tower.
66
+ sees_images: bool
67
+
68
+ def render(self, system: str, user: str, prefill: str, *, open_ended: bool = False) -> str:
69
+ """Wrap the message bodies in the host's chat template and append the prefill.
70
+
71
+ open_ended returns only the span up to the end of the user body, which is
72
+ exactly what every branch of one request shares. The closed render of a state
73
+ must begin with the open_ended render of that same state character for
74
+ character, and encode must split at that same point, so that
75
+ encode(prefix) + encode(suffix) equals encode(prefix + suffix) there.
76
+ """
77
+ ...
78
+
79
+ def encode(self, text: str) -> list[int]:
80
+ """Token ids for a fragment, with no special tokens added."""
81
+ ...
82
+
83
+ def encode_prefix(self, text: str, image: Any = None) -> tuple[list[int], dict[str, Any]]:
84
+ """Prefix token ids, plus whatever tensors its forward pass needs for the image."""
85
+ ...
86
+
87
+ def score(
88
+ self,
89
+ prefix_ids: list[int],
90
+ suffix_ids: list[list[int]],
91
+ label_counts: list[int],
92
+ vision: dict[str, Any] | None = None,
93
+ ) -> list[BranchLogits]:
94
+ """Read the label logits at each branch's final position, sharing one prefix.
95
+
96
+ The returned list carries one entry per suffix, in the order given. Entry i's
97
+ z is exactly label_counts[i] wide, ordered to match LABEL_ALPHABET[:count],
98
+ since the caller maps those positions straight onto the branch's options.
99
+ """
100
+ ...
101
+
102
+
103
+ def verify_backend(backend: Backend) -> list[int]:
104
+ """Prove a backend meets the port at load, and return its label ids.
105
+
106
+ A backend that gets any of this wrong returns wrong probabilities rather than
107
+ raising, so it is worth one call before a node goes live. The render check
108
+ mirrors what _suffix_ids does per branch on every request.
109
+ """
110
+ if not isinstance(getattr(backend, "model_id", None), str):
111
+ raise BackendContractError(
112
+ f"{type(backend).__name__} declares no model_id, so the fitted temperature "
113
+ f"cannot be looked up and the answer would be silently miscalibrated"
114
+ )
115
+
116
+ open_ended = backend.render(_PROBE_SYSTEM, _PROBE_STATE, ANSWER_PREFILL, open_ended=True)
117
+ closed = backend.render(_PROBE_SYSTEM, _PROBE_STATE + _PROBE_SUFFIX, ANSWER_PREFILL)
118
+
119
+ if not closed.startswith(open_ended):
120
+ raise BackendContractError(
121
+ f"{type(backend).__name__}.render did not begin its closed render with the "
122
+ f"open-ended render of the same state, so every branch would be encoded at "
123
+ f"the wrong offset"
124
+ )
125
+ return verify_label_ids(backend, closed, MAX_LABELS_PER_BRANCH)
@@ -0,0 +1,412 @@
1
+ """The transformers backend: model loading and the single-forward-pass branch scorer.
2
+
3
+ No token is ever generated. Every probability comes from the logit row at one
4
+ position, the position the answer prefill forces to be the answer.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import copy
10
+ import os
11
+ import threading
12
+ import warnings
13
+ from collections.abc import Iterator
14
+ from contextlib import contextmanager
15
+ from typing import Any
16
+
17
+ # cuBLAS needs a fixed workspace before torch initialises CUDA to keep its
18
+ # reductions reproducible.
19
+ os.environ.setdefault("CUBLAS_WORKSPACE_CONFIG", ":4096:8")
20
+
21
+ import torch
22
+ from torch.nn.attention import SDPBackend, sdpa_kernel
23
+ from transformers import (
24
+ AutoConfig,
25
+ AutoModelForCausalLM,
26
+ AutoModelForImageTextToText,
27
+ AutoProcessor,
28
+ AutoTokenizer,
29
+ )
30
+ from transformers.cache_utils import DynamicCache
31
+
32
+ from ..config import Config, canonical_model_id
33
+ from .base import BranchLogits, VisionUnsupportedError, verify_backend
34
+
35
+ # Expanding the shared prefix across batch rows costs this much per row, so the
36
+ # batch has to shrink as the context grows.
37
+ KV_BUDGET_BYTES = 6 * 1024**3
38
+
39
+ # A chunk's suffix tokens, rows times padded width, checked only when a branch widens
40
+ # the chunk it joins. The KV budget bounds what the prefix cache costs, and this bounds
41
+ # what a wide branch can charge the narrow ones beside it. Anything from 512 to 2048
42
+ # measured the same, and 4096 upward is clearly worse. ab_branch_packing.py.
43
+ CHUNK_TOKEN_CEILING = 2048
44
+
45
+ # torch 2.11 ships no FlashAttention kernel on Windows. Left to choose for itself the
46
+ # dispatcher then reaches the math backend, which builds the full attention matrix. On an
47
+ # 8k prefill that measured 147 seconds and 27.6 GB against 1.3 seconds and 9.2 GB here.
48
+ PREFERRED_ATTENTION = (SDPBackend.CUDNN_ATTENTION, SDPBackend.EFFICIENT_ATTENTION)
49
+
50
+ # When no preferred backend probes clean, the host's own enable flags would otherwise
51
+ # decide which kernel runs, so two hosts could answer one request differently. Pinning
52
+ # this order makes the choice a function of the hardware. Math is last and always works.
53
+ FALLBACK_ATTENTION = (
54
+ SDPBackend.FLASH_ATTENTION,
55
+ SDPBackend.CUDNN_ATTENTION,
56
+ SDPBackend.EFFICIENT_ATTENTION,
57
+ SDPBackend.MATH,
58
+ )
59
+
60
+
61
+ def _usable_attention_backends(device: torch.device, dtype: torch.dtype) -> tuple[SDPBackend, ...]:
62
+ """Probe which preferred backends have a kernel for this dtype on this build.
63
+
64
+ A backend that serves bfloat16 can be missing for another dtype, so the probe has to
65
+ run at the dtype the model will use.
66
+ """
67
+ usable: list[SDPBackend] = []
68
+
69
+ if device.type != "cuda":
70
+ return ()
71
+ # The probe draws from its own generator so that constructing a backend does not
72
+ # shift the host's global RNG stream.
73
+ generator = torch.Generator(device=device)
74
+ query = torch.randn(1, 4, 64, 64, device=device, dtype=dtype, generator=generator)
75
+ key = torch.randn(1, 2, 64, 64, device=device, dtype=dtype, generator=generator)
76
+ mask = torch.zeros(1, 1, 64, 64, device=device, dtype=dtype)
77
+ attention = torch.nn.functional.scaled_dot_product_attention
78
+ for backend in PREFERRED_ATTENTION:
79
+ try:
80
+ # A rejected backend warns on the way out, which is the answer we came for.
81
+ with warnings.catch_warnings(), sdpa_kernel(backend):
82
+ warnings.simplefilter("ignore", UserWarning)
83
+ attention(query, key, key, is_causal=True, enable_gqa=True)
84
+ attention(query, key, key, attn_mask=mask, enable_gqa=True)
85
+ usable.append(backend)
86
+ except RuntimeError:
87
+ continue
88
+ return tuple(usable)
89
+
90
+
91
+ # Two windows open on two threads would interleave their saves, and the second to
92
+ # exit would restore the pinned values rather than the host's. The service already
93
+ # serialises on its own GPU lock, which is always taken before this one.
94
+ _WINDOW = threading.RLock()
95
+
96
+
97
+ @contextmanager
98
+ def _determinism() -> Iterator[None]:
99
+ """Hold the torch globals that decide bit-exact reductions, for one forward pass.
100
+
101
+ Every alternative value buys speed by giving up bit-exactness, so none is tunable.
102
+ torch exposes none of them as a call argument, so scoping them means setting them
103
+ here and putting the host's values back after. A ComfyUI host sharing this process
104
+ keeps its own settings everywhere outside the block.
105
+
106
+ On the CUDA attention path only allow_bf16_reduced_precision_reduction moves a logit
107
+ on either shipped model. The rest are kept because cudnn picks convolution algorithms
108
+ by timing, which is specific to the card, and `ab_determinism_scope.py` measured one.
109
+
110
+ The sdp and fp16 accumulation settings are held for a third reason, that ComfyUI turns
111
+ each of them on and neither is reachable by that sweep. `ab_math_sdp_reduction.py`
112
+ measures the sdp one on the math backend.
113
+
114
+ Where torch has per-backend matmul slots, the coarse getter raises once a host has
115
+ set a slot directly, so only the raw slots are saved, pinned and put back there.
116
+ """
117
+ with _WINDOW:
118
+ saved_benchmark = torch.backends.cudnn.benchmark
119
+ saved_deterministic = torch.backends.cudnn.deterministic
120
+ saved_bf16 = torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction
121
+ # torch ships no stub for the mkldnn matmul slot, so it is reached by name.
122
+ mkldnn: Any = getattr(torch.backends.mkldnn, "matmul", None)
123
+ has_slots = hasattr(torch.backends.cuda.matmul, "fp32_precision") and hasattr(
124
+ mkldnn, "fp32_precision"
125
+ )
126
+ saved_matmul = None if has_slots else torch.get_float32_matmul_precision()
127
+ saved_cuda = torch.backends.cuda.matmul.fp32_precision if has_slots else None
128
+ saved_mkldnn = mkldnn.fp32_precision if has_slots else None
129
+ # ComfyUI turns this on at import, at comfy/model_management.py:569. It governs
130
+ # the math attention backend, which runs only when attention_backends is empty.
131
+ has_sdp = hasattr(torch.backends.cuda, "allow_fp16_bf16_reduction_math_sdp") and hasattr(
132
+ torch.backends.cuda, "fp16_bf16_reduction_math_sdp_allowed"
133
+ )
134
+ saved_sdp = bool(
135
+ has_sdp and torch.backends.cuda.fp16_bf16_reduction_math_sdp_allowed()
136
+ )
137
+ # A bare --fast turns this on, since comfy/cli_args.py then enables every
138
+ # PerformanceFeature. It is the fp16 sibling of the bf16 reduction above.
139
+ has_fp16_acc = hasattr(torch.backends.cuda.matmul, "allow_fp16_accumulation")
140
+ saved_fp16_acc = bool(
141
+ has_fp16_acc and torch.backends.cuda.matmul.allow_fp16_accumulation
142
+ )
143
+
144
+ torch.backends.cudnn.benchmark = False
145
+ torch.backends.cudnn.deterministic = True
146
+ torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction = False
147
+ if has_slots:
148
+ torch.backends.cuda.matmul.fp32_precision = "ieee"
149
+ mkldnn.fp32_precision = "ieee"
150
+ else:
151
+ torch.set_float32_matmul_precision("highest")
152
+ if has_sdp:
153
+ torch.backends.cuda.allow_fp16_bf16_reduction_math_sdp(False)
154
+ if has_fp16_acc:
155
+ torch.backends.cuda.matmul.allow_fp16_accumulation = False
156
+ try:
157
+ yield
158
+ finally:
159
+ torch.backends.cudnn.benchmark = saved_benchmark
160
+ torch.backends.cudnn.deterministic = saved_deterministic
161
+ torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction = saved_bf16
162
+ if has_slots:
163
+ torch.backends.cuda.matmul.fp32_precision = saved_cuda
164
+ mkldnn.fp32_precision = saved_mkldnn
165
+ elif saved_matmul is not None:
166
+ torch.set_float32_matmul_precision(saved_matmul)
167
+ if has_sdp:
168
+ torch.backends.cuda.allow_fp16_bf16_reduction_math_sdp(saved_sdp)
169
+ if has_fp16_acc:
170
+ torch.backends.cuda.matmul.allow_fp16_accumulation = saved_fp16_acc
171
+
172
+
173
+ class HFBackend:
174
+ """Reads logits off a Hugging Face model held in this process."""
175
+
176
+ def __init__(self, config: Config) -> None:
177
+ self.config = config
178
+ self.model_id = config.model_id
179
+ self.canonical_model_id = canonical_model_id(config.model_id)
180
+ # None is what transformers already means by "use the HF cache", so one keyword
181
+ # covers both the project folder and the global location with no branch here.
182
+ cache_dir = str(config.models_dir) if config.models_dir else None
183
+
184
+ loaded = AutoConfig.from_pretrained(config.model_id, cache_dir=cache_dir)
185
+ self.sees_images = hasattr(loaded, "vision_config")
186
+ self.processor = None
187
+
188
+ if self.sees_images:
189
+ self.processor = AutoProcessor.from_pretrained(config.model_id, cache_dir=cache_dir)
190
+ self.tokenizer = self.processor.tokenizer
191
+ self.model = AutoModelForImageTextToText.from_pretrained(
192
+ config.model_id,
193
+ cache_dir=cache_dir,
194
+ dtype=getattr(torch, config.dtype),
195
+ device_map=config.device,
196
+ ).eval()
197
+ else:
198
+ self.tokenizer = AutoTokenizer.from_pretrained(config.model_id, cache_dir=cache_dir)
199
+ self.model = AutoModelForCausalLM.from_pretrained(
200
+ config.model_id,
201
+ cache_dir=cache_dir,
202
+ dtype=getattr(torch, config.dtype),
203
+ device_map=config.device,
204
+ ).eval()
205
+
206
+ self.label_ids = verify_backend(self)
207
+ self._label_id_tensor = torch.tensor(self.label_ids, device=self.model.device)
208
+ self._pad_id = self.tokenizer.pad_token_id or self.tokenizer.eos_token_id
209
+ self._kv_bytes_per_token = self._measure_kv_bytes_per_token()
210
+ self.attention_backends = _usable_attention_backends(self.model.device, self.model.dtype)
211
+
212
+ @contextmanager
213
+ def _pinned_forward(self) -> Iterator[None]:
214
+ """Establish the environment every forward pass in this class needs.
215
+
216
+ The attention backend and the determinism globals are both process wide, so
217
+ each pass sets them and puts them back rather than pinning them at load.
218
+ Every self.model call belongs inside this block.
219
+ """
220
+ with _determinism():
221
+ if not self.attention_backends:
222
+ with sdpa_kernel(list(FALLBACK_ATTENTION), set_priority=True):
223
+ yield
224
+ return
225
+ with sdpa_kernel(list(self.attention_backends)):
226
+ yield
227
+
228
+ def _measure_kv_bytes_per_token(self) -> int:
229
+ cfg = getattr(self.model.config, "text_config", self.model.config)
230
+ heads = getattr(cfg, "num_key_value_heads", cfg.num_attention_heads)
231
+ head_dim = getattr(cfg, "head_dim", cfg.hidden_size // cfg.num_attention_heads)
232
+ element = torch.empty((), dtype=getattr(torch, self.config.dtype)).element_size()
233
+ return int(2 * cfg.num_hidden_layers * heads * head_dim * element)
234
+
235
+ def render(self, system: str, user: str, prefill: str, *, open_ended: bool = False) -> str:
236
+ """Apply the chat template, then append the answer prefill.
237
+
238
+ open_ended returns only the user-message body, before the template closes the
239
+ turn, which is exactly the span every branch shares.
240
+ """
241
+ messages = [{"role": "system", "content": system}, {"role": "user", "content": user}]
242
+ rendered: str = self.tokenizer.apply_chat_template(
243
+ messages, tokenize=False, add_generation_prompt=True
244
+ )
245
+ if not open_ended:
246
+ return rendered + prefill
247
+ return rendered[: rendered.rindex(user) + len(user)]
248
+
249
+ def encode(self, text: str) -> list[int]:
250
+ ids: list[int] = self.tokenizer.encode(text, add_special_tokens=False)
251
+ return ids
252
+
253
+ def encode_prefix(self, text: str, image: Any = None) -> tuple[list[int], dict[str, Any]]:
254
+ """Prefix token ids, plus the vision tensors its forward pass needs.
255
+
256
+ The processor expands the single image marker into one token per patch, so
257
+ the count is decided here rather than guessed. Only the prefix carries an
258
+ image, which is what keeps every branch a plain text suffix.
259
+ """
260
+ if image is None:
261
+ return self.encode(text), {}
262
+ if not self.sees_images or self.processor is None:
263
+ raise VisionUnsupportedError(
264
+ f"{self.config.model_id} has no vision tower, so it cannot read an image"
265
+ )
266
+ batch = self.processor(text=[text], images=[image], return_tensors="pt",
267
+ add_special_tokens=False)
268
+ ids: list[int] = batch["input_ids"][0].tolist()
269
+ vision = {k: v.to(self.model.device) for k, v in batch.items()
270
+ if k in ("pixel_values", "image_grid_thw", "mm_token_type_ids")}
271
+ return ids, vision
272
+
273
+ def _rows_per_chunk(self, prefix_len: int, suffix_len: int) -> int:
274
+ if not self.config.batch_branches:
275
+ return 1
276
+ per_row = self._kv_bytes_per_token * (prefix_len + suffix_len)
277
+ affordable = max(1, KV_BUDGET_BYTES // max(per_row, 1))
278
+ return int(min(self.config.max_batch_rows, affordable))
279
+
280
+ def _pack_chunks(self, suffix_ids: list[list[int]], prefix_len: int) -> list[list[int]]:
281
+ """Group branch indices into chunks of similar suffix length.
282
+
283
+ Every row in a chunk is left-padded to the chunk's longest suffix, and each
284
+ padded token costs a full forward plus attention over the whole prefix. Taking
285
+ branches in request order drags short ones to the longest width, which measured
286
+ 86.7 percent waste on a mixed request. The sort is stable, so equal lengths keep
287
+ request order and the grouping stays a pure function of the request.
288
+ """
289
+ if not self.config.batch_branches:
290
+ return [[index] for index in range(len(suffix_ids))]
291
+
292
+ chunks: list[list[int]] = []
293
+ current: list[int] = []
294
+
295
+ for index in sorted(range(len(suffix_ids)), key=lambda i: len(suffix_ids[i])):
296
+ width = len(suffix_ids[index])
297
+ allowed = self._rows_per_chunk(prefix_len, width)
298
+ rows = len(current) + 1
299
+ # The ceiling exists to stop one wide branch dragging narrow ones out to
300
+ # its width. A candidate no wider than the chunk adds no padding, so only
301
+ # the row cap applies and a request of one shape chunks as it always did.
302
+ widens = bool(current) and width > len(suffix_ids[current[-1]])
303
+ over_ceiling = widens and rows * width > CHUNK_TOKEN_CEILING
304
+ if current and (rows > allowed or over_ceiling):
305
+ chunks.append(current)
306
+ current = [index]
307
+ else:
308
+ current.append(index)
309
+ if current:
310
+ chunks.append(current)
311
+ return chunks
312
+
313
+ @torch.inference_mode()
314
+ def score(
315
+ self, prefix_ids: list[int], suffix_ids: list[list[int]], label_counts: list[int],
316
+ vision: dict[str, Any] | None = None,
317
+ ) -> list[BranchLogits]:
318
+ """Score every branch against one shared prefix.
319
+
320
+ The prefix is encoded once. Each chunk then copies that cache, broadcasts
321
+ it across the chunk's rows, and reads the final position of every row in
322
+ a single forward pass.
323
+ """
324
+ device = self.model.device
325
+ scored: dict[int, BranchLogits] = {}
326
+ rope_delta: torch.Tensor | None = None
327
+
328
+ if not suffix_ids:
329
+ return []
330
+
331
+ prefix_tensor = torch.tensor([prefix_ids], device=device)
332
+ seed = DynamicCache()
333
+ with self._pinned_forward():
334
+ self.model(
335
+ input_ids=prefix_tensor,
336
+ attention_mask=torch.ones_like(prefix_tensor),
337
+ past_key_values=seed,
338
+ use_cache=True,
339
+ logits_to_keep=1,
340
+ **(vision or {}),
341
+ )
342
+ # An image makes positions three dimensional and shifts every later token, so
343
+ # the branches reuse the offset this pass wrote before another pass overwrites it.
344
+ if vision:
345
+ rope_delta = getattr(self.model.model, "rope_deltas", None)
346
+ if rope_delta is not None:
347
+ rope_delta = rope_delta.clone()
348
+
349
+ # The port promises one row per suffix in the order given, so the packed
350
+ # chunks are scattered back rather than concatenated.
351
+ for chunk in self._pack_chunks(suffix_ids, len(prefix_ids)):
352
+ rows = self._score_chunk(
353
+ prefix_ids, seed,
354
+ [suffix_ids[i] for i in chunk], [label_counts[i] for i in chunk], rope_delta,
355
+ )
356
+ for slot, row in zip(chunk, rows, strict=True):
357
+ scored[slot] = row
358
+ return [scored[index] for index in range(len(suffix_ids))]
359
+
360
+ def _score_chunk(
361
+ self,
362
+ prefix_ids: list[int],
363
+ seed: DynamicCache,
364
+ suffix_ids: list[list[int]],
365
+ label_counts: list[int],
366
+ rope_delta: torch.Tensor | None = None,
367
+ ) -> list[BranchLogits]:
368
+ device = self.model.device
369
+ rows = len(suffix_ids)
370
+ prefix_len = len(prefix_ids)
371
+ width = max(len(s) for s in suffix_ids)
372
+ results: list[BranchLogits] = []
373
+
374
+ # Left-padding puts every row's final real token at the same index, so a
375
+ # single kept position serves the whole batch.
376
+ input_ids = torch.full((rows, width), self._pad_id, dtype=torch.long)
377
+ attention = torch.zeros((rows, prefix_len + width), dtype=torch.long)
378
+ attention[:, :prefix_len] = 1
379
+ positions = torch.zeros((rows, width), dtype=torch.long)
380
+ for row, suffix in enumerate(suffix_ids):
381
+ pad = width - len(suffix)
382
+ input_ids[row, pad:] = torch.tensor(suffix)
383
+ attention[row, prefix_len + pad :] = 1
384
+ positions[row, pad:] = torch.arange(prefix_len, prefix_len + len(suffix))
385
+
386
+ placed = positions.to(device)
387
+ if rope_delta is not None:
388
+ placed = (placed + rope_delta.to(device)).unsqueeze(0).expand(3, rows, width)
389
+
390
+ cache = copy.deepcopy(seed)
391
+ cache.batch_repeat_interleave(rows)
392
+ with self._pinned_forward():
393
+ output = self.model(
394
+ input_ids=input_ids.to(device),
395
+ attention_mask=attention.to(device),
396
+ position_ids=placed,
397
+ past_key_values=cache,
398
+ use_cache=True,
399
+ logits_to_keep=1,
400
+ )
401
+ final = output.logits[:, -1, :].float()
402
+
403
+ # Index on device so only the handful of label logits crosses the bus.
404
+ selected = final.index_select(1, self._label_id_tensor)
405
+ full_norm = torch.logsumexp(final, dim=-1)
406
+ del cache, output, final
407
+
408
+ for row, count in enumerate(label_counts):
409
+ z = selected[row, :count]
410
+ mass = float(torch.exp(torch.logsumexp(z, dim=-1) - full_norm[row]))
411
+ results.append(BranchLogits(z=z.double().cpu().numpy(), candidate_mass=mass))
412
+ return results