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,141 @@
1
+ """Batch calibration: a running per-label prior learned from served traffic.
2
+
3
+ The label prior is token bias, the model's standing preference for the token
4
+ "A" over "B". It is estimated as a running mean of the uncalibrated label
5
+ distribution and subtracted in log space before the softmax. No labelled data
6
+ and no extra forward passes are required.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import contextlib
12
+ import json
13
+ import math
14
+ import threading
15
+ from pathlib import Path
16
+ from uuid import uuid4
17
+
18
+ import numpy as np
19
+
20
+ from .config import PRIOR_MIN_OBSERVATIONS, PRIOR_SMOOTHING
21
+
22
+
23
+ def _stored_mean(key: str, vector: object) -> list[float] | None:
24
+ """Return the mean of one `kind:width` bucket, or None when the stored value is unusable."""
25
+ kind, separator, width = key.partition(":")
26
+
27
+ if not kind or separator != ":" or not (width.isascii() and width.isdigit()):
28
+ return None
29
+ if not isinstance(vector, list) or len(vector) != int(width):
30
+ return None
31
+ if not all(isinstance(value, int | float) and not isinstance(value, bool) for value in vector):
32
+ return None
33
+ values = [float(value) for value in vector]
34
+ if not all(math.isfinite(value) and value >= 0.0 for value in values):
35
+ return None
36
+ return values
37
+
38
+
39
+ class PriorStore:
40
+ def __init__(
41
+ self, path: Path | None, fingerprint: str, min_observations: int = PRIOR_MIN_OBSERVATIONS
42
+ ) -> None:
43
+ self.path = path
44
+ self.fingerprint = fingerprint
45
+ self.min_observations = min_observations
46
+ self._lock = threading.Lock()
47
+ self._means: dict[str, list[float]] = {}
48
+ self._counts: dict[str, int] = {}
49
+ self._load()
50
+
51
+ @staticmethod
52
+ def _bucket(kind: str, label_count: int) -> str:
53
+ return f"{kind}:{label_count}"
54
+
55
+ def _load(self) -> None:
56
+ means: dict[str, list[float]] = {}
57
+ counts: dict[str, int] = {}
58
+
59
+ if self.path is None or not self.path.exists():
60
+ return
61
+ try:
62
+ payload = json.loads(self.path.read_text(encoding="utf-8"))
63
+ except (OSError, json.JSONDecodeError):
64
+ return
65
+ if not isinstance(payload, dict):
66
+ return
67
+ # A prior measured under a different model or prompt describes a different
68
+ # distribution, so a mismatch starts over rather than blending the two.
69
+ if payload.get("fingerprint") != self.fingerprint:
70
+ return
71
+ stored_means = payload.get("means")
72
+ stored_counts = payload.get("counts")
73
+ if not isinstance(stored_means, dict) or not isinstance(stored_counts, dict):
74
+ return
75
+
76
+ # The path is user-visible state, so a hand-edited bucket is a reachable input.
77
+ # One that does not match its own key starts over the way a mismatch does, rather
78
+ # than reaching observe and raising on every later request.
79
+ for key, vector in stored_means.items():
80
+ mean = _stored_mean(key, vector) if isinstance(key, str) else None
81
+ count = stored_counts.get(key)
82
+ if mean is None or not isinstance(count, int) or isinstance(count, bool) or count < 0:
83
+ continue
84
+ means[key] = mean
85
+ counts[key] = count
86
+
87
+ self._means = means
88
+ self._counts = counts
89
+
90
+ def save(self) -> None:
91
+ # A path with no final component, `.` or a drive root, has no staged sibling to
92
+ # name, and with_name raises ValueError, which the OSError suppression misses.
93
+ if self.path is None or not self.path.name:
94
+ return
95
+ payload = {"fingerprint": self.fingerprint, "means": self._means, "counts": self._counts}
96
+
97
+ # Two processes can hold a store over one path, and a truncating write hands a
98
+ # reader a partial file that loads as an empty prior. A rename is atomic.
99
+ with self._lock:
100
+ staged = self.path.with_name(f"{self.path.name}.{uuid4().hex}.tmp")
101
+ try:
102
+ staged.write_text(json.dumps(payload, indent=2, sort_keys=True), encoding="utf-8")
103
+ staged.replace(self.path)
104
+ except OSError:
105
+ # A rename that keeps failing would otherwise leave one staged file per
106
+ # request, since the uuid name is never reused.
107
+ with contextlib.suppress(OSError):
108
+ staged.unlink(missing_ok=True)
109
+
110
+ def observe(self, kind: str, probabilities: np.ndarray) -> None:
111
+ """Fold one uncalibrated distribution into the running mean."""
112
+ bucket = self._bucket(kind, len(probabilities))
113
+
114
+ with self._lock:
115
+ count = self._counts.get(bucket, 0)
116
+ mean = np.asarray(self._means.get(bucket, [0.0] * len(probabilities)), dtype=np.float64)
117
+ self._means[bucket] = ((count * mean + probabilities) / (count + 1)).tolist()
118
+ self._counts[bucket] = count + 1
119
+
120
+ def log_prior(self, kind: str, label_count: int) -> np.ndarray | None:
121
+ """Return the log-prior to subtract, or None while the estimate is thin."""
122
+ bucket = self._bucket(kind, label_count)
123
+
124
+ with self._lock:
125
+ count = self._counts.get(bucket, 0)
126
+ mean = self._means.get(bucket)
127
+ if mean is None or count < self.min_observations or len(mean) != label_count:
128
+ return None
129
+ prior = np.asarray(mean, dtype=np.float64)
130
+ if not np.all(prior >= 0) or prior.sum() <= 0:
131
+ return None
132
+ prior = prior / prior.sum()
133
+ prior = (1.0 - PRIOR_SMOOTHING) * prior + PRIOR_SMOOTHING / label_count
134
+ # Centring keeps the subtraction from shifting overall scale, which only
135
+ # the temperature should control.
136
+ log_prior = np.log(prior)
137
+ return log_prior - log_prior.mean()
138
+
139
+ def stats(self) -> dict[str, int]:
140
+ with self._lock:
141
+ return dict(self._counts)
@@ -0,0 +1,329 @@
1
+ """Pipeline coordinator: request in, Jev-shaped answers out."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import hashlib
6
+ import random
7
+ from dataclasses import dataclass, field, replace
8
+
9
+ import numpy as np
10
+
11
+ from .backends.base import Backend, BackendContractError, BranchLogits
12
+ from .calibrate import PriorStore
13
+ from .config import ANSWER_PREFILL, Config, fitted_temperature
14
+ from .deps import MissingDependencyError
15
+ from .prompt import (
16
+ PROMPT_VERSION,
17
+ SYSTEM_PROMPT,
18
+ Branch,
19
+ branch_content,
20
+ build_branches,
21
+ prefix_content,
22
+ )
23
+ from .schema import (
24
+ Answer,
25
+ ChoiceAnswer,
26
+ ChoiceQuestion,
27
+ NoulAnswer,
28
+ NoulQuestion,
29
+ Question,
30
+ ScoreAnswer,
31
+ ScoreQuestion,
32
+ SystemOneRequest,
33
+ SystemOneResponse,
34
+ Usage,
35
+ )
36
+ from .scoring import (
37
+ choice_confidence,
38
+ combine_escape,
39
+ expected_score,
40
+ normalise_levels,
41
+ restricted_softmax,
42
+ score_confidence,
43
+ )
44
+ from .vision import extract_image
45
+
46
+ # The label prior is a property of the token and how many labels compete, so
47
+ # branches sharing a shape share a bucket. A score's letter names a fixed rung, so a
48
+ # score never shares a bucket with a choice of the same width.
49
+ _PRIOR_KIND = {
50
+ "choice": "choice",
51
+ "member": "choice",
52
+ "score_level": "binary",
53
+ "score_joint": "score",
54
+ "noul": "binary",
55
+ }
56
+
57
+
58
+ @dataclass
59
+ class Diagnostics:
60
+ """Per-branch signals that Jev does not return, surfaced on request."""
61
+
62
+ candidate_mass: dict[str, list[float]] = field(default_factory=dict)
63
+ prior_applied: dict[str, bool] = field(default_factory=dict)
64
+ branch_counts: dict[str, int] = field(default_factory=dict)
65
+
66
+
67
+ def load_model(model_id: str | None = None, config: Config | None = None) -> Backend:
68
+ """Load a model through transformers and return it as a Backend.
69
+
70
+ `model_id` names a Hugging Face repo or a local directory, and overrides the one on
71
+ `config`. Where the weights land is `config.models_dir`. A ComfyUI node skips this
72
+ entirely and passes its own Backend over a model the workflow already loaded.
73
+
74
+ The import sits inside the call so the core install needs no torch.
75
+ """
76
+ config = config if config is not None else Config()
77
+ if model_id is not None:
78
+ config = replace(config, model_id=model_id)
79
+
80
+ try:
81
+ from .backends.hf import HFBackend
82
+ except ImportError as error:
83
+ raise MissingDependencyError(
84
+ "no backend was passed and the local transformers backend is not installed. "
85
+ 'Install it with: pip install "logit-classifier[hf]"'
86
+ ) from error
87
+ return HFBackend(config)
88
+
89
+
90
+ def _model_identity(backend: Backend) -> str:
91
+ """Return the id the fitted temperature and the prior fingerprint key on."""
92
+ return getattr(backend, "canonical_model_id", None) or backend.model_id
93
+
94
+
95
+ class Classifier:
96
+ """Turns one System One request into one calibrated answer per question.
97
+
98
+ A caller that already holds a loaded model, such as a ComfyUI node over a
99
+ resident text encoder, passes its own backend and no weights are loaded here.
100
+
101
+ One instance belongs to one backend. The fitted temperature and the prior bucket
102
+ are both derived from the backend's model at construction, so swapping the
103
+ backend afterwards leaves both belonging to the previous model.
104
+ """
105
+
106
+ def __init__(self, config: Config, backend: Backend | None = None) -> None:
107
+ self.config = config
108
+ self.backend = backend if backend is not None else load_model(config=config)
109
+ # Keyed on the model the backend actually loaded, not on a Config field a
110
+ # caller passing its own backend never had reason to set.
111
+ self.temperature = (
112
+ config.temperature if config.temperature is not None
113
+ else fitted_temperature(_model_identity(self.backend))
114
+ )
115
+ self.priors = PriorStore(config.calibration_path, self._fingerprint())
116
+
117
+ def _fingerprint(self) -> str:
118
+ """Identify the prompt shape a stored prior was measured under."""
119
+ parts = [_model_identity(self.backend), ANSWER_PREFILL, SYSTEM_PROMPT, self.config.score_method,
120
+ str(PROMPT_VERSION), str(self.config.abstain)]
121
+ return hashlib.sha256("|".join(parts).encode("utf-8")).hexdigest()[:16]
122
+
123
+ def persist_prior(self) -> None:
124
+ """Write the running prior, unless the caller turned the prior off.
125
+
126
+ The service calls this again at shutdown, so the gate lives here rather than
127
+ at each call site.
128
+ """
129
+ if self.config.use_prior_debias:
130
+ self.priors.save()
131
+
132
+ def _calibrated(self, branch: Branch, logits: BranchLogits) -> tuple[np.ndarray, bool]:
133
+ """Raw distribution feeds the prior estimate, calibrated one is returned."""
134
+ kind = _PRIOR_KIND[branch.kind]
135
+ # The escape label means "none of these" wherever it sits, so its mass is content
136
+ # rather than letter bias, and only the lettered options get a prior.
137
+ lettered = branch.label_count - 1 if branch.kind == "member" else branch.label_count
138
+ letters_prior: np.ndarray | None = None
139
+ log_prior: np.ndarray | None = None
140
+
141
+ # One switch turns learning, application and persistence off together, so a host
142
+ # that asked for no debias gets the same numbers on every call.
143
+ if self.config.use_prior_debias:
144
+ self.priors.observe(kind, restricted_softmax(logits.z[:lettered]))
145
+ letters_prior = self.priors.log_prior(kind, lettered)
146
+ if letters_prior is not None:
147
+ log_prior = np.zeros(branch.label_count, dtype=np.float64)
148
+ log_prior[:lettered] = letters_prior
149
+ calibrated = restricted_softmax(logits.z, log_prior, self.temperature)
150
+ return calibrated, log_prior is not None
151
+
152
+ def _letterings(self, questions: dict[str, Question]) -> list[dict[str, Question]]:
153
+ """One question map per lettering, the first in the order the caller gave.
154
+
155
+ Only choice options are reordered. A score lists its levels lowest first, so
156
+ that order carries meaning, and a noul has two fixed sides. Each lettering
157
+ draws from a fixed seed, so the same request still answers the same way.
158
+ """
159
+ rounds = [questions]
160
+
161
+ for seed in range(1, max(1, self.config.permutations)):
162
+ shuffler = random.Random(seed)
163
+ reordered: dict[str, Question] = {}
164
+ for qid, question in questions.items():
165
+ if not isinstance(question, ChoiceQuestion):
166
+ continue
167
+ names = list(question.criteria)
168
+ shuffler.shuffle(names)
169
+ criteria = {name: question.criteria[name] for name in names}
170
+ reordered[qid] = replace(question, criteria=criteria)
171
+ if reordered:
172
+ rounds.append(reordered)
173
+ return rounds
174
+
175
+ def _suffix_ids(self, state: object, branches: list[Branch], has_image: bool,
176
+ prefix_text: str) -> list[list[int]]:
177
+ """Each branch's tokens after the shared prefix, which it must begin with."""
178
+ ids: list[list[int]] = []
179
+
180
+ for branch in branches:
181
+ rendered = self.backend.render(
182
+ SYSTEM_PROMPT, branch_content(state, branch, has_image), ANSWER_PREFILL
183
+ )
184
+ if not rendered.startswith(prefix_text):
185
+ raise BackendContractError(
186
+ f"{type(self.backend).__name__}.render did not begin the closed render "
187
+ f"of question {branch.question_id!r} with the open-ended render of the "
188
+ "same state, so the shared prefix cannot be split off by offset"
189
+ )
190
+ ids.append(self.backend.encode(rendered[len(prefix_text):]))
191
+ return ids
192
+
193
+ def classify(
194
+ self, request: SystemOneRequest, *, allow_image_paths: bool = True
195
+ ) -> tuple[SystemOneResponse, Diagnostics]:
196
+ rounds = self._letterings(request.questions)
197
+ diagnostics = Diagnostics()
198
+ probabilities: list[np.ndarray] = []
199
+ backend_name = type(self.backend).__name__
200
+ drafts: list[dict[str, Answer]] = []
201
+ start = 0
202
+
203
+ state, image = extract_image(request.state, allow_paths=allow_image_paths)
204
+ seen = image is not None
205
+ prefix_text = self.backend.render(
206
+ SYSTEM_PROMPT, prefix_content(state, seen), ANSWER_PREFILL, open_ended=True
207
+ )
208
+ prefix_ids, vision = self.backend.encode_prefix(prefix_text, image)
209
+ round_branches = [build_branches(q, self.config.score_method, self.config.abstain)
210
+ for q in rounds]
211
+ branches: list[Branch] = [b for lettering in round_branches for b in lettering]
212
+ suffix_ids = self._suffix_ids(state, branches, seen, prefix_text)
213
+ scored = self.backend.score(prefix_ids, suffix_ids,
214
+ [b.label_count for b in branches], vision)
215
+
216
+ if len(scored) != len(branches):
217
+ raise BackendContractError(
218
+ f"{backend_name}.score returned {len(scored)} rows for "
219
+ f"{len(branches)} branches"
220
+ )
221
+ for branch, logits in zip(branches, scored, strict=True):
222
+ if len(logits.z) != branch.label_count:
223
+ raise BackendContractError(
224
+ f"{backend_name}.score returned a {len(logits.z)}-wide row for "
225
+ f"question {branch.question_id}, which needs {branch.label_count}"
226
+ )
227
+ calibrated, applied = self._calibrated(branch, logits)
228
+ probabilities.append(calibrated)
229
+ diagnostics.candidate_mass.setdefault(branch.question_id, []).append(
230
+ round(logits.candidate_mass, 6)
231
+ )
232
+ diagnostics.prior_applied[branch.question_id] = applied
233
+ diagnostics.branch_counts[branch.question_id] = (
234
+ diagnostics.branch_counts.get(branch.question_id, 0) + 1
235
+ )
236
+ self.persist_prior()
237
+
238
+ for questions, lettering in zip(rounds, round_branches, strict=True):
239
+ window = slice(start, start + len(lettering))
240
+ start += len(lettering)
241
+ drafts.append(self._round_answers(questions, lettering, probabilities[window]))
242
+ answers = self._merge(request.questions, drafts)
243
+
244
+ usage = Usage(
245
+ input_tokens=len(prefix_ids) + sum(len(s) for s in suffix_ids),
246
+ output_tokens=len(branches),
247
+ )
248
+ response = SystemOneResponse(
249
+ model=self.config.served_model_id, answers=answers, usage=usage
250
+ )
251
+ return response, diagnostics
252
+
253
+ def _round_answers(self, questions: dict[str, Question], branches: list[Branch],
254
+ probabilities: list[np.ndarray]) -> dict[str, Answer]:
255
+ answers: dict[str, Answer] = {}
256
+
257
+ for qid, question in questions.items():
258
+ span = [i for i, b in enumerate(branches) if b.question_id == qid]
259
+ answers[qid] = self._assemble(question, [branches[i] for i in span],
260
+ [probabilities[i] for i in span])
261
+ return answers
262
+
263
+ def _merge(self, questions: dict[str, Question],
264
+ drafts: list[dict[str, Answer]]) -> dict[str, Answer]:
265
+ """Average each option's probability across the letterings that scored it."""
266
+ answers: dict[str, Answer] = {}
267
+
268
+ for qid in questions:
269
+ parts = [draft[qid] for draft in drafts if qid in draft]
270
+ if len(parts) == 1:
271
+ answers[qid] = parts[0]
272
+ continue
273
+ answers[qid] = self._average_choice([p for p in parts if isinstance(p, ChoiceAnswer)])
274
+ return answers
275
+
276
+ def _average_choice(self, parts: list[ChoiceAnswer]) -> ChoiceAnswer:
277
+ names = list(parts[0].probabilities)
278
+ stacked = np.array([[part.probabilities[name] for name in names] for part in parts])
279
+ mean = stacked.mean(axis=0)
280
+ mean = mean / mean.sum()
281
+ declined = [part.abstain for part in parts if part.abstain is not None]
282
+ return ChoiceAnswer(
283
+ choice=names[int(np.argmax(mean))],
284
+ confidence=round(choice_confidence(mean), 6),
285
+ probabilities={name: round(float(p), 6) for name, p in zip(names, mean, strict=True)},
286
+ abstain=round(float(np.mean(declined)), 6) if declined else None,
287
+ )
288
+
289
+ def _assemble(self, question: Question, branches: list[Branch],
290
+ probabilities: list[np.ndarray]) -> Answer:
291
+ if isinstance(question, ChoiceQuestion):
292
+ return self._choice_answer(question, branches, probabilities)
293
+ if isinstance(question, ScoreQuestion):
294
+ return self._score_answer(question, branches, probabilities)
295
+ if isinstance(question, NoulQuestion):
296
+ return NoulAnswer(noul=round(float(probabilities[0][0]), 6))
297
+ raise TypeError(f"unsupported question type {type(question)!r}")
298
+
299
+ def _choice_answer(self, question: ChoiceQuestion, branches: list[Branch],
300
+ probabilities: list[np.ndarray]) -> ChoiceAnswer:
301
+ names = list(question.criteria)
302
+ declined: float | None = None
303
+
304
+ if branches[0].kind == "choice":
305
+ flat = probabilities[0]
306
+ else:
307
+ flat, declined = combine_escape(probabilities)
308
+ mapping = {name: round(float(p), 6) for name, p in zip(names, flat, strict=True)}
309
+ winner = names[int(np.argmax(flat))]
310
+ return ChoiceAnswer(
311
+ choice=winner,
312
+ confidence=round(choice_confidence(flat), 6),
313
+ probabilities=mapping,
314
+ abstain=None if declined is None else round(declined, 6),
315
+ )
316
+
317
+ def _score_answer(self, question: ScoreQuestion, branches: list[Branch],
318
+ probabilities: list[np.ndarray]) -> ScoreAnswer:
319
+ if branches[0].kind == "score_joint":
320
+ levels = probabilities[0]
321
+ else:
322
+ # Each level was judged alone, so index 0 of each branch is its yes.
323
+ levels = normalise_levels(np.array([float(p[0]) for p in probabilities]))
324
+ return ScoreAnswer(
325
+ score=round(expected_score(levels), 6),
326
+ confidence=round(score_confidence(levels), 6),
327
+ legend={str(i): level for i, level in enumerate(question.criteria)},
328
+ probabilities={str(i): round(float(p), 6) for i, p in enumerate(levels)},
329
+ )
@@ -0,0 +1,51 @@
1
+ """The `logit-classifier` command.
2
+
3
+ `serve` is the front door for anyone who installed from PyPI rather than cloning,
4
+ since it needs no uvicorn invocation and no import path.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import argparse
10
+
11
+ from . import __version__
12
+ from .config import Config
13
+ from .deps import require
14
+
15
+
16
+ def _serve(host: str, port: int, reload: bool) -> int:
17
+ # uvicorn imports the app from a string, so a missing fastapi surfaces as a raw
18
+ # traceback out of its importer rather than as the line that installs the extra.
19
+ require("fastapi", "service")
20
+ uvicorn = require("uvicorn", "service")
21
+ uvicorn.run("logit_classifier.service:app", host=host, port=port, reload=reload)
22
+ return 0
23
+
24
+
25
+ def _show_config() -> int:
26
+ config = Config.from_env()
27
+ for name, value in vars(config).items():
28
+ print(f"{name}={value}")
29
+ return 0
30
+
31
+
32
+ def _parser(prog: str) -> argparse.ArgumentParser:
33
+ parser = argparse.ArgumentParser(prog=prog, description=__doc__.splitlines()[0])
34
+ parser.add_argument("--version", action="version", version=__version__)
35
+ commands = parser.add_subparsers(dest="command", required=True)
36
+
37
+ serve = commands.add_parser("serve", help="run the HTTP service and the browser page")
38
+ serve.add_argument("--host", default="127.0.0.1")
39
+ serve.add_argument("--port", type=int, default=8077)
40
+ serve.add_argument("--reload", action="store_true", help="restart on a source change")
41
+
42
+ commands.add_parser("config", help="print the configuration the LOGIT_ variables produce")
43
+ return parser
44
+
45
+
46
+ def main(argv: list[str] | None = None, prog: str = "logit-classifier") -> int:
47
+ args = _parser(prog).parse_args(argv)
48
+
49
+ if args.command == "serve":
50
+ return _serve(args.host, args.port, args.reload)
51
+ return _show_config()