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.
- logit_classifier/__init__.py +90 -0
- logit_classifier/__main__.py +10 -0
- logit_classifier/backends/__init__.py +5 -0
- logit_classifier/backends/base.py +125 -0
- logit_classifier/backends/hf.py +412 -0
- logit_classifier/calibrate.py +141 -0
- logit_classifier/classifier.py +329 -0
- logit_classifier/cli.py +51 -0
- logit_classifier/config.py +187 -0
- logit_classifier/deps.py +30 -0
- logit_classifier/errors.py +16 -0
- logit_classifier/labels.py +91 -0
- logit_classifier/prompt.py +154 -0
- logit_classifier/py.typed +0 -0
- logit_classifier/schema.py +285 -0
- logit_classifier/scoring.py +102 -0
- logit_classifier/service.py +174 -0
- logit_classifier/vision.py +95 -0
- logit_classifier/web/index.html +411 -0
- logit_classifier-0.1.0.dist-info/METADATA +354 -0
- logit_classifier-0.1.0.dist-info/RECORD +24 -0
- logit_classifier-0.1.0.dist-info/WHEEL +4 -0
- logit_classifier-0.1.0.dist-info/entry_points.txt +2 -0
- logit_classifier-0.1.0.dist-info/licenses/LICENSE +674 -0
|
@@ -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,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
|