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,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
|
+
)
|
logit_classifier/cli.py
ADDED
|
@@ -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()
|