logit-classifier 0.1.0__tar.gz → 0.2.0__tar.gz
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-0.1.0 → logit_classifier-0.2.0}/PKG-INFO +23 -4
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/README.md +22 -3
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/pyproject.toml +9 -2
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/src/logit_classifier/__init__.py +3 -2
- logit_classifier-0.2.0/src/logit_classifier/_version.py +3 -0
- logit_classifier-0.2.0/src/logit_classifier/backends/_torch_window.py +96 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/src/logit_classifier/backends/base.py +4 -0
- logit_classifier-0.2.0/src/logit_classifier/backends/comfy_clip.py +414 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/src/logit_classifier/backends/hf.py +1 -83
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/src/logit_classifier/classifier.py +33 -3
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/src/logit_classifier/config.py +2 -1
- logit_classifier-0.2.0/src/logit_classifier/tags.py +121 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/src/logit_classifier/vision.py +21 -7
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests/test_unit.py +988 -5
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/.gitignore +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/FINDINGS.md +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/LICENSE +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/examples/compare_models.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/examples/http_client.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/examples/images.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/examples/own_backend.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/examples/question_types.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/examples/quickstart.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/src/logit_classifier/__main__.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/src/logit_classifier/backends/__init__.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/src/logit_classifier/calibrate.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/src/logit_classifier/cli.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/src/logit_classifier/deps.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/src/logit_classifier/errors.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/src/logit_classifier/labels.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/src/logit_classifier/prompt.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/src/logit_classifier/py.typed +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/src/logit_classifier/schema.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/src/logit_classifier/scoring.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/src/logit_classifier/service.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/src/logit_classifier/web/index.html +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests/fixtures/banking77_test.json +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests/fixtures/eval_set.json +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests/fixtures/many_options_request.json +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests/fixtures/quickstart_request.json +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests/test_model.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests-AB/ab_branch_packing.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests-AB/ab_determinism_scope.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests-AB/ab_env.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests-AB/ab_math_sdp_reduction.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests-AB/ab_multi_label.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests-AB/ab_noul_wording.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests-AB/ab_temperature.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests-AB/banking77.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests-AB/benchmark.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests-AB/evaluate.py +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests-AB/inputs/_full.png +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests-AB/inputs/_sheet.png +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests-AB/inputs/low-L.png +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests-AB/inputs/low-R.png +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests-AB/inputs/mid-L.png +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests-AB/inputs/mid-R.png +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests-AB/inputs/top-L.png +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests-AB/inputs/top-R.png +0 -0
- {logit_classifier-0.1.0 → logit_classifier-0.2.0}/tests-AB/tune_groups.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.5
|
|
2
2
|
Name: logit-classifier
|
|
3
|
-
Version: 0.
|
|
3
|
+
Version: 0.2.0
|
|
4
4
|
Summary: Local zero-shot classifier for text and images. Declares options, returns a calibrated probability for each, generates no text. Accepts the TypeSafe System One request format.
|
|
5
5
|
Project-URL: Repository, https://github.com/Blakeem/logit-classifier
|
|
6
6
|
Project-URL: Bug Tracker, https://github.com/Blakeem/logit-classifier/issues
|
|
@@ -169,7 +169,7 @@ the model.
|
|
|
169
169
|
|
|
170
170
|
```json
|
|
171
171
|
{
|
|
172
|
-
"model": "logit-classifier-0.
|
|
172
|
+
"model": "logit-classifier-0.2.0",
|
|
173
173
|
"answers": {
|
|
174
174
|
"department": {
|
|
175
175
|
"type": "choice",
|
|
@@ -210,7 +210,7 @@ only in process.
|
|
|
210
210
|
"questions": { "sky": { "type": "noul", "instructions": "This crop shows the night sky" } } }
|
|
211
211
|
```
|
|
212
212
|
|
|
213
|
-
The image is encoded once for the entire request
|
|
213
|
+
The image is encoded once for the entire request. So asking several questions about one
|
|
214
214
|
picture costs little more than asking one.
|
|
215
215
|
|
|
216
216
|
A model with no vision tower rejects the request rather than ignoring the picture.
|
|
@@ -271,7 +271,8 @@ Batch composition and padding length both change the low bits of a bfloat16 forw
|
|
|
271
271
|
and both are pure functions of the request. So the same question asked inside two
|
|
272
272
|
different requests can differ slightly. Set `LOGIT_BATCH_BRANCHES=0` to score each branch
|
|
273
273
|
alone, which makes a question independent of the questions sent with it and costs one
|
|
274
|
-
forward pass per branch.
|
|
274
|
+
forward pass per branch. `ComfyClipBackend` reads no environment variable, so it takes
|
|
275
|
+
`batch_branches=False` as a keyword instead.
|
|
275
276
|
|
|
276
277
|
A forward pass needs several process-global torch settings held at known values. Torch
|
|
277
278
|
exposes none of them as a call argument, so the backend sets them around each pass and
|
|
@@ -320,6 +321,24 @@ the render and label contracts and returns one token id per label.
|
|
|
320
321
|
`logit_classifier.backends.hf.HFBackend` is the transformers implementation to read
|
|
321
322
|
against.
|
|
322
323
|
|
|
324
|
+
`ComfyClipBackend` is the backend over a Qwen3-VL text encoder that the workflow already
|
|
325
|
+
loaded, such as the one Krea 2 uses. So no second model is loaded into VRAM.
|
|
326
|
+
|
|
327
|
+
```python
|
|
328
|
+
from logit_classifier import Classifier, Config, NoulQuestion, SystemOneRequest
|
|
329
|
+
from logit_classifier.backends.comfy_clip import ComfyClipBackend
|
|
330
|
+
|
|
331
|
+
classifier = Classifier(Config(), backend=ComfyClipBackend(clip))
|
|
332
|
+
question = NoulQuestion(instructions="This image visibly contains a dragon")
|
|
333
|
+
request = SystemOneRequest(state="", questions={"dragon": question})
|
|
334
|
+
response, _ = classifier.classify(request, image=image)
|
|
335
|
+
```
|
|
336
|
+
|
|
337
|
+
The `image` keyword takes an image the host already decoded, such as a ComfyUI IMAGE of
|
|
338
|
+
shape `[1, H, W, 3]`. An empty state asks about the image alone. Questions share one forward
|
|
339
|
+
pass of up to 4,096 tokens, and the image counts toward that limit. Any text encoder other
|
|
340
|
+
than Qwen3-VL raises `UnsupportedModelError`.
|
|
341
|
+
|
|
323
342
|
## Differences From Jev
|
|
324
343
|
|
|
325
344
|
This project matches the System One request and response format and the documented limits.
|
|
@@ -136,7 +136,7 @@ the model.
|
|
|
136
136
|
|
|
137
137
|
```json
|
|
138
138
|
{
|
|
139
|
-
"model": "logit-classifier-0.
|
|
139
|
+
"model": "logit-classifier-0.2.0",
|
|
140
140
|
"answers": {
|
|
141
141
|
"department": {
|
|
142
142
|
"type": "choice",
|
|
@@ -177,7 +177,7 @@ only in process.
|
|
|
177
177
|
"questions": { "sky": { "type": "noul", "instructions": "This crop shows the night sky" } } }
|
|
178
178
|
```
|
|
179
179
|
|
|
180
|
-
The image is encoded once for the entire request
|
|
180
|
+
The image is encoded once for the entire request. So asking several questions about one
|
|
181
181
|
picture costs little more than asking one.
|
|
182
182
|
|
|
183
183
|
A model with no vision tower rejects the request rather than ignoring the picture.
|
|
@@ -238,7 +238,8 @@ Batch composition and padding length both change the low bits of a bfloat16 forw
|
|
|
238
238
|
and both are pure functions of the request. So the same question asked inside two
|
|
239
239
|
different requests can differ slightly. Set `LOGIT_BATCH_BRANCHES=0` to score each branch
|
|
240
240
|
alone, which makes a question independent of the questions sent with it and costs one
|
|
241
|
-
forward pass per branch.
|
|
241
|
+
forward pass per branch. `ComfyClipBackend` reads no environment variable, so it takes
|
|
242
|
+
`batch_branches=False` as a keyword instead.
|
|
242
243
|
|
|
243
244
|
A forward pass needs several process-global torch settings held at known values. Torch
|
|
244
245
|
exposes none of them as a call argument, so the backend sets them around each pass and
|
|
@@ -287,6 +288,24 @@ the render and label contracts and returns one token id per label.
|
|
|
287
288
|
`logit_classifier.backends.hf.HFBackend` is the transformers implementation to read
|
|
288
289
|
against.
|
|
289
290
|
|
|
291
|
+
`ComfyClipBackend` is the backend over a Qwen3-VL text encoder that the workflow already
|
|
292
|
+
loaded, such as the one Krea 2 uses. So no second model is loaded into VRAM.
|
|
293
|
+
|
|
294
|
+
```python
|
|
295
|
+
from logit_classifier import Classifier, Config, NoulQuestion, SystemOneRequest
|
|
296
|
+
from logit_classifier.backends.comfy_clip import ComfyClipBackend
|
|
297
|
+
|
|
298
|
+
classifier = Classifier(Config(), backend=ComfyClipBackend(clip))
|
|
299
|
+
question = NoulQuestion(instructions="This image visibly contains a dragon")
|
|
300
|
+
request = SystemOneRequest(state="", questions={"dragon": question})
|
|
301
|
+
response, _ = classifier.classify(request, image=image)
|
|
302
|
+
```
|
|
303
|
+
|
|
304
|
+
The `image` keyword takes an image the host already decoded, such as a ComfyUI IMAGE of
|
|
305
|
+
shape `[1, H, W, 3]`. An empty state asks about the image alone. Questions share one forward
|
|
306
|
+
pass of up to 4,096 tokens, and the image counts toward that limit. Any text encoder other
|
|
307
|
+
than Qwen3-VL raises `UnsupportedModelError`.
|
|
308
|
+
|
|
290
309
|
## Differences From Jev
|
|
291
310
|
|
|
292
311
|
This project matches the System One request and response format and the documented limits.
|
|
@@ -136,7 +136,8 @@ show_error_codes = true
|
|
|
136
136
|
[[tool.mypy.overrides]]
|
|
137
137
|
# transformers publishes no complete stubs, and accelerate none at all. torch is
|
|
138
138
|
# listed so the type gate still runs where the hf extra is absent, such as CI.
|
|
139
|
-
|
|
139
|
+
# The ComfyUI host supplies comfy, and neither this venv nor CI has it.
|
|
140
|
+
module = ["transformers.*", "accelerate.*", "torch.*", "comfy.*"]
|
|
140
141
|
ignore_missing_imports = true
|
|
141
142
|
|
|
142
143
|
[[tool.mypy.overrides]]
|
|
@@ -148,6 +149,12 @@ module = ["logit_classifier.backends.hf"]
|
|
|
148
149
|
disallow_untyped_calls = false
|
|
149
150
|
disallow_untyped_decorators = false
|
|
150
151
|
|
|
152
|
+
[[tool.mypy.overrides]]
|
|
153
|
+
# torch leaves fp16_bf16_reduction_math_sdp_allowed unannotated, and a line ignore
|
|
154
|
+
# would go unused, and so fail, where the hf extra is absent and torch is Any.
|
|
155
|
+
module = ["logit_classifier.backends._torch_window"]
|
|
156
|
+
disallow_untyped_calls = false
|
|
157
|
+
|
|
151
158
|
[tool.pytest.ini_options]
|
|
152
159
|
pythonpath = ["src"]
|
|
153
160
|
testpaths = ["tests"]
|
|
@@ -159,7 +166,7 @@ requires = ["hatchling"]
|
|
|
159
166
|
build-backend = "hatchling.build"
|
|
160
167
|
|
|
161
168
|
[tool.hatch.version]
|
|
162
|
-
path = "src/logit_classifier/
|
|
169
|
+
path = "src/logit_classifier/_version.py"
|
|
163
170
|
|
|
164
171
|
[tool.hatch.build.targets.wheel]
|
|
165
172
|
packages = ["src/logit_classifier"]
|
|
@@ -9,10 +9,12 @@ with the `[hf]` extra and the HTTP service with `[service]`.
|
|
|
9
9
|
|
|
10
10
|
from __future__ import annotations
|
|
11
11
|
|
|
12
|
+
from ._version import __version__
|
|
12
13
|
from .backends.base import (
|
|
13
14
|
Backend,
|
|
14
15
|
BackendContractError,
|
|
15
16
|
BranchLogits,
|
|
17
|
+
UnsupportedModelError,
|
|
16
18
|
VisionUnsupportedError,
|
|
17
19
|
verify_backend,
|
|
18
20
|
)
|
|
@@ -45,8 +47,6 @@ from .schema import (
|
|
|
45
47
|
)
|
|
46
48
|
from .vision import ImageError
|
|
47
49
|
|
|
48
|
-
__version__ = "0.1.0"
|
|
49
|
-
|
|
50
50
|
# The ComfyUI socket type a node pack declares for a loaded classifier. It lives
|
|
51
51
|
# here so the library and the packs cannot drift apart on the spelling.
|
|
52
52
|
COMFY_SOCKET_TYPE = "LOGIT_CLASSIFIER"
|
|
@@ -79,6 +79,7 @@ __all__ = [
|
|
|
79
79
|
"ScoreQuestion",
|
|
80
80
|
"SystemOneRequest",
|
|
81
81
|
"SystemOneResponse",
|
|
82
|
+
"UnsupportedModelError",
|
|
82
83
|
"Usage",
|
|
83
84
|
"VisionUnsupportedError",
|
|
84
85
|
"__version__",
|
|
@@ -0,0 +1,96 @@
|
|
|
1
|
+
"""The determinism window every backend's forward pass runs inside.
|
|
2
|
+
|
|
3
|
+
torch is imported inside the window rather than at module scope, so a backend over a
|
|
4
|
+
host's own model can import this without pulling torch into a core-only process.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
import threading
|
|
10
|
+
from collections.abc import Iterator
|
|
11
|
+
from contextlib import contextmanager
|
|
12
|
+
from typing import Any
|
|
13
|
+
|
|
14
|
+
# Two windows open on two threads would interleave their saves, and the second to
|
|
15
|
+
# exit would restore the pinned values rather than the host's. The service already
|
|
16
|
+
# serialises on its own GPU lock, which is always taken before this one.
|
|
17
|
+
_WINDOW = threading.RLock()
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
@contextmanager
|
|
21
|
+
def _determinism() -> Iterator[None]:
|
|
22
|
+
"""Hold the torch globals that decide bit-exact reductions, for one forward pass.
|
|
23
|
+
|
|
24
|
+
Every alternative value buys speed by giving up bit-exactness, so none is tunable.
|
|
25
|
+
torch exposes none of them as a call argument, so scoping them means setting them
|
|
26
|
+
here and putting the host's values back after. A ComfyUI host sharing this process
|
|
27
|
+
keeps its own settings everywhere outside the block.
|
|
28
|
+
|
|
29
|
+
On the CUDA attention path only allow_bf16_reduced_precision_reduction moves a logit
|
|
30
|
+
on either shipped model. The rest are kept because cudnn picks convolution algorithms
|
|
31
|
+
by timing, which is specific to the card, and `ab_determinism_scope.py` measured one.
|
|
32
|
+
|
|
33
|
+
The sdp and fp16 accumulation settings are held for a third reason, that ComfyUI turns
|
|
34
|
+
each of them on and neither is reachable by that sweep. `ab_math_sdp_reduction.py`
|
|
35
|
+
measures the sdp one on the math backend.
|
|
36
|
+
|
|
37
|
+
Where torch has per-backend matmul slots, the coarse getter raises once a host has
|
|
38
|
+
set a slot directly, so only the raw slots are saved, pinned and put back there.
|
|
39
|
+
"""
|
|
40
|
+
import torch
|
|
41
|
+
|
|
42
|
+
with _WINDOW:
|
|
43
|
+
saved_benchmark = torch.backends.cudnn.benchmark
|
|
44
|
+
saved_deterministic = torch.backends.cudnn.deterministic
|
|
45
|
+
saved_bf16 = torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction
|
|
46
|
+
# torch ships no stub for the mkldnn matmul slot, so it is reached by name.
|
|
47
|
+
mkldnn: Any = getattr(torch.backends.mkldnn, "matmul", None)
|
|
48
|
+
has_slots = hasattr(torch.backends.cuda.matmul, "fp32_precision") and hasattr(
|
|
49
|
+
mkldnn, "fp32_precision"
|
|
50
|
+
)
|
|
51
|
+
saved_matmul = None if has_slots else torch.get_float32_matmul_precision()
|
|
52
|
+
saved_cuda = torch.backends.cuda.matmul.fp32_precision if has_slots else None
|
|
53
|
+
saved_mkldnn = mkldnn.fp32_precision if has_slots else None
|
|
54
|
+
# ComfyUI turns this on at import, at comfy/model_management.py:569. It governs
|
|
55
|
+
# the math attention backend, which runs only when attention_backends is empty.
|
|
56
|
+
has_sdp = hasattr(torch.backends.cuda, "allow_fp16_bf16_reduction_math_sdp") and hasattr(
|
|
57
|
+
torch.backends.cuda, "fp16_bf16_reduction_math_sdp_allowed"
|
|
58
|
+
)
|
|
59
|
+
saved_sdp = bool(
|
|
60
|
+
has_sdp and torch.backends.cuda.fp16_bf16_reduction_math_sdp_allowed()
|
|
61
|
+
)
|
|
62
|
+
# A bare --fast turns this on, since comfy/cli_args.py then enables every
|
|
63
|
+
# PerformanceFeature. It is the fp16 sibling of the bf16 reduction above.
|
|
64
|
+
has_fp16_acc = hasattr(torch.backends.cuda.matmul, "allow_fp16_accumulation")
|
|
65
|
+
saved_fp16_acc = bool(
|
|
66
|
+
has_fp16_acc and torch.backends.cuda.matmul.allow_fp16_accumulation
|
|
67
|
+
)
|
|
68
|
+
|
|
69
|
+
# Pinning sits inside the try, so a setter that raises part way still restores the host.
|
|
70
|
+
try:
|
|
71
|
+
torch.backends.cudnn.benchmark = False
|
|
72
|
+
torch.backends.cudnn.deterministic = True
|
|
73
|
+
torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction = False
|
|
74
|
+
if has_slots:
|
|
75
|
+
torch.backends.cuda.matmul.fp32_precision = "ieee"
|
|
76
|
+
mkldnn.fp32_precision = "ieee"
|
|
77
|
+
else:
|
|
78
|
+
torch.set_float32_matmul_precision("highest")
|
|
79
|
+
if has_sdp:
|
|
80
|
+
torch.backends.cuda.allow_fp16_bf16_reduction_math_sdp(False)
|
|
81
|
+
if has_fp16_acc:
|
|
82
|
+
torch.backends.cuda.matmul.allow_fp16_accumulation = False
|
|
83
|
+
yield
|
|
84
|
+
finally:
|
|
85
|
+
torch.backends.cudnn.benchmark = saved_benchmark
|
|
86
|
+
torch.backends.cudnn.deterministic = saved_deterministic
|
|
87
|
+
torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction = saved_bf16
|
|
88
|
+
if has_slots:
|
|
89
|
+
torch.backends.cuda.matmul.fp32_precision = saved_cuda
|
|
90
|
+
mkldnn.fp32_precision = saved_mkldnn
|
|
91
|
+
elif saved_matmul is not None:
|
|
92
|
+
torch.set_float32_matmul_precision(saved_matmul)
|
|
93
|
+
if has_sdp:
|
|
94
|
+
torch.backends.cuda.allow_fp16_bf16_reduction_math_sdp(saved_sdp)
|
|
95
|
+
if has_fp16_acc:
|
|
96
|
+
torch.backends.cuda.matmul.allow_fp16_accumulation = saved_fp16_acc
|
|
@@ -32,6 +32,10 @@ class BackendContractError(LogitClassifierError, RuntimeError):
|
|
|
32
32
|
"""A backend returned rows the port does not allow."""
|
|
33
33
|
|
|
34
34
|
|
|
35
|
+
class UnsupportedModelError(LogitClassifierError, ValueError):
|
|
36
|
+
"""A host passed a loaded model this backend cannot read a logit row from."""
|
|
37
|
+
|
|
38
|
+
|
|
35
39
|
@dataclass(frozen=True)
|
|
36
40
|
class BranchLogits:
|
|
37
41
|
"""Raw label logits for one branch, before any calibration."""
|