jev-compatible-server 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.
- jev_compatible_server/__init__.py +5 -0
- jev_compatible_server/app.py +144 -0
- jev_compatible_server/backends.py +320 -0
- jev_compatible_server/batching.py +68 -0
- jev_compatible_server/bosun.py +229 -0
- jev_compatible_server/causal_options.py +228 -0
- jev_compatible_server/classifier_adapters.py +472 -0
- jev_compatible_server/cross_encoder.py +71 -0
- jev_compatible_server/custom_heads.py +704 -0
- jev_compatible_server/encoder_decoder.py +630 -0
- jev_compatible_server/gliner2.py +40 -0
- jev_compatible_server/hidden_state_probe.py +384 -0
- jev_compatible_server/laya.py +135 -0
- jev_compatible_server/native_systemone.py +248 -0
- jev_compatible_server/protocol.py +104 -0
- jev_compatible_server/public-models.json +981 -0
- jev_compatible_server/registry.py +283 -0
- jev_compatible_server/runtime.py +237 -0
- jev_compatible_server/sequence_classifier.py +219 -0
- jev_compatible_server-0.1.0.dist-info/METADATA +157 -0
- jev_compatible_server-0.1.0.dist-info/RECORD +23 -0
- jev_compatible_server-0.1.0.dist-info/WHEEL +4 -0
- jev_compatible_server-0.1.0.dist-info/entry_points.txt +2 -0
|
@@ -0,0 +1,472 @@
|
|
|
1
|
+
"""Classifier readouts for NLI and GLiClass decision checkpoints.
|
|
2
|
+
|
|
3
|
+
Both adapters retain the server's single candidate representation: Jev question
|
|
4
|
+
candidates are compiled once, then projected into the checkpoint's native
|
|
5
|
+
classifier output space before the shared answer aggregation runs.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import json
|
|
11
|
+
import math
|
|
12
|
+
from collections.abc import Mapping, Sequence
|
|
13
|
+
from dataclasses import dataclass
|
|
14
|
+
from pathlib import Path
|
|
15
|
+
from typing import Any
|
|
16
|
+
|
|
17
|
+
from .encoder_decoder import (
|
|
18
|
+
MarginTask,
|
|
19
|
+
_mapping,
|
|
20
|
+
_template,
|
|
21
|
+
aggregate_margin_answers,
|
|
22
|
+
compile_margin_tasks,
|
|
23
|
+
decision_metadata,
|
|
24
|
+
normalized_entropy_confidence,
|
|
25
|
+
render_content,
|
|
26
|
+
)
|
|
27
|
+
from .protocol import (
|
|
28
|
+
ChoiceAnswer,
|
|
29
|
+
ChoiceQuestion,
|
|
30
|
+
DecisionRequest,
|
|
31
|
+
DecisionResponse,
|
|
32
|
+
NoulAnswer,
|
|
33
|
+
NoulQuestion,
|
|
34
|
+
ScoreAnswer,
|
|
35
|
+
ScoreQuestion,
|
|
36
|
+
Usage,
|
|
37
|
+
)
|
|
38
|
+
from .runtime import DecisionRuntime, RuntimeErrorBase, softmax
|
|
39
|
+
|
|
40
|
+
VERDICT_ABSTENTION_ID = "__insufficient_evidence__"
|
|
41
|
+
VERDICT_ABSTENTION_LABEL = "insufficient evidence"
|
|
42
|
+
VERDICT_MAX_SUBSTANTIVE_CANDIDATES = 24
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def _positive_int(value: Any, name: str, default: int) -> int:
|
|
46
|
+
resolved = default if value is None else value
|
|
47
|
+
if not isinstance(resolved, int) or isinstance(resolved, bool) or resolved <= 0:
|
|
48
|
+
raise RuntimeErrorBase(f"{name} must be a positive integer")
|
|
49
|
+
return resolved
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def _device(torch: Any, requested: Any, default: str) -> str:
|
|
53
|
+
value = default if requested is None else requested
|
|
54
|
+
if not isinstance(value, str):
|
|
55
|
+
raise RuntimeErrorBase("decision.device must be a string")
|
|
56
|
+
target = "cuda" if value == "auto" and torch.cuda.is_available() else "cpu" if value == "auto" else value
|
|
57
|
+
if target.startswith("cuda") and not torch.cuda.is_available():
|
|
58
|
+
raise RuntimeErrorBase("CUDA was requested, but no CUDA device is available")
|
|
59
|
+
return target
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def _loader_kwargs(loader: Mapping[str, Any]) -> dict[str, Any]:
|
|
63
|
+
trust_remote_code = loader.get("trust_remote_code", False)
|
|
64
|
+
if not isinstance(trust_remote_code, bool):
|
|
65
|
+
raise RuntimeErrorBase("decision.loader.trust_remote_code must be boolean")
|
|
66
|
+
kwargs: dict[str, Any] = {"trust_remote_code": trust_remote_code}
|
|
67
|
+
revision = loader.get("revision")
|
|
68
|
+
if revision is not None:
|
|
69
|
+
if not isinstance(revision, str):
|
|
70
|
+
raise RuntimeErrorBase("decision.loader.revision must be a string")
|
|
71
|
+
kwargs["revision"] = revision
|
|
72
|
+
return kwargs
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def project_entailment_logits(
|
|
76
|
+
logits: Any, label2id: Mapping[str, Any], entailment_label: str = "entailment"
|
|
77
|
+
) -> Any:
|
|
78
|
+
"""Select one NLI entailment logit per candidate without model-name heuristics."""
|
|
79
|
+
|
|
80
|
+
if not isinstance(entailment_label, str) or not entailment_label:
|
|
81
|
+
raise RuntimeErrorBase("decision.entailment_label must be a non-empty string")
|
|
82
|
+
matches = [
|
|
83
|
+
value
|
|
84
|
+
for label, value in label2id.items()
|
|
85
|
+
if isinstance(label, str) and label.casefold() == entailment_label.casefold()
|
|
86
|
+
]
|
|
87
|
+
if len(matches) != 1 or not isinstance(matches[0], int) or isinstance(matches[0], bool):
|
|
88
|
+
raise RuntimeErrorBase(
|
|
89
|
+
f"NLI model must define exactly one {entailment_label!r} label in label2id"
|
|
90
|
+
)
|
|
91
|
+
index = matches[0]
|
|
92
|
+
if getattr(logits, "ndim", None) != 2 or index < 0 or index >= logits.shape[1]:
|
|
93
|
+
raise RuntimeErrorBase("NLI model output does not contain the configured entailment logit")
|
|
94
|
+
values = logits[:, index]
|
|
95
|
+
if not bool(values.isfinite().all()):
|
|
96
|
+
raise RuntimeErrorBase("NLI model produced non-finite entailment logits")
|
|
97
|
+
return values
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
class NLIEntailmentBackend(DecisionRuntime):
|
|
101
|
+
"""Project each candidate pair through an NLI head's entailment dimension."""
|
|
102
|
+
|
|
103
|
+
def __init__(
|
|
104
|
+
self, model_id: str, *, config: dict[str, Any] | None = None, device: str = "auto"
|
|
105
|
+
) -> None:
|
|
106
|
+
try:
|
|
107
|
+
import torch
|
|
108
|
+
from transformers import AutoModelForSequenceClassification, AutoTokenizer
|
|
109
|
+
except ImportError as exc: # pragma: no cover - optional dependency
|
|
110
|
+
raise RuntimeErrorBase("NLIEntailmentBackend requires transformers and torch") from exc
|
|
111
|
+
|
|
112
|
+
self.model_name = str((config or {}).get("model", model_id))
|
|
113
|
+
self.config = config or {}
|
|
114
|
+
self.metadata = decision_metadata(self.config)
|
|
115
|
+
if self.metadata.get("readout") != "nli_entailment":
|
|
116
|
+
raise RuntimeErrorBase("NLIEntailmentBackend requires decision.readout=nli_entailment")
|
|
117
|
+
self._torch = torch
|
|
118
|
+
self._batch_size = _positive_int(self.metadata.get("batch_size"), "decision.batch_size", 32)
|
|
119
|
+
input_config = _mapping(self.metadata.get("input"), "decision.input")
|
|
120
|
+
self._max_length = _positive_int(input_config.get("max_length"), "decision.input.max_length", 512)
|
|
121
|
+
self._premise_template = _template(input_config.get("premise_template"), "decision.input.premise_template")
|
|
122
|
+
self._hypothesis_template = _template(input_config.get("hypothesis_template"), "decision.input.hypothesis_template")
|
|
123
|
+
|
|
124
|
+
loader = _mapping(self.metadata.get("loader", {}), "decision.loader")
|
|
125
|
+
base_model = loader.get("base_model", model_id)
|
|
126
|
+
tokenizer_id = loader.get("tokenizer", base_model)
|
|
127
|
+
if not isinstance(base_model, str) or not isinstance(tokenizer_id, str):
|
|
128
|
+
raise RuntimeErrorBase("decision.loader base_model and tokenizer must be strings")
|
|
129
|
+
kwargs = _loader_kwargs(loader)
|
|
130
|
+
self._tokenizer = AutoTokenizer.from_pretrained(tokenizer_id, **kwargs)
|
|
131
|
+
self._model = AutoModelForSequenceClassification.from_pretrained(base_model, **kwargs)
|
|
132
|
+
self._model.to(_device(torch, self.metadata.get("device"), device))
|
|
133
|
+
self._model.eval()
|
|
134
|
+
self._input_device = next(self._model.parameters()).device
|
|
135
|
+
self._label2id = self._model.config.label2id
|
|
136
|
+
self._entailment_label = self.metadata.get("entailment_label", "entailment")
|
|
137
|
+
|
|
138
|
+
def _pairs(self, tasks: Sequence[MarginTask]) -> tuple[list[str], list[str]]:
|
|
139
|
+
premises: list[str] = []
|
|
140
|
+
hypotheses: list[str] = []
|
|
141
|
+
for task in tasks:
|
|
142
|
+
try:
|
|
143
|
+
premises.append(self._premise_template.format(state=task.query))
|
|
144
|
+
hypotheses.append(
|
|
145
|
+
self._hypothesis_template.format(
|
|
146
|
+
instructions=task.instruction, candidate=task.document
|
|
147
|
+
)
|
|
148
|
+
)
|
|
149
|
+
except KeyError as exc:
|
|
150
|
+
raise RuntimeErrorBase(
|
|
151
|
+
f"NLI input template references an unknown field: {exc.args[0]}"
|
|
152
|
+
) from exc
|
|
153
|
+
return premises, hypotheses
|
|
154
|
+
|
|
155
|
+
def decide_batch(self, requests: Sequence[DecisionRequest]) -> list[DecisionResponse]:
|
|
156
|
+
compiled = [compile_margin_tasks(request, self.metadata) for request in requests]
|
|
157
|
+
tasks = [task for request_tasks in compiled for task in request_tasks]
|
|
158
|
+
premises, hypotheses = self._pairs(tasks)
|
|
159
|
+
margins: list[float] = []
|
|
160
|
+
token_counts: list[int] = []
|
|
161
|
+
for start in range(0, len(tasks), self._batch_size):
|
|
162
|
+
encoded = self._tokenizer(
|
|
163
|
+
premises[start : start + self._batch_size],
|
|
164
|
+
hypotheses[start : start + self._batch_size],
|
|
165
|
+
padding=True,
|
|
166
|
+
truncation=True,
|
|
167
|
+
max_length=self._max_length,
|
|
168
|
+
return_tensors="pt",
|
|
169
|
+
)
|
|
170
|
+
encoded = {name: value.to(self._input_device) for name, value in encoded.items()}
|
|
171
|
+
with self._torch.inference_mode():
|
|
172
|
+
output = self._model(**encoded)
|
|
173
|
+
margins.extend(
|
|
174
|
+
float(value)
|
|
175
|
+
for value in project_entailment_logits(
|
|
176
|
+
output.logits.float(), self._label2id, self._entailment_label
|
|
177
|
+
).cpu().tolist()
|
|
178
|
+
)
|
|
179
|
+
token_counts.extend(int(value) for value in encoded["attention_mask"].sum(dim=1).cpu().tolist())
|
|
180
|
+
|
|
181
|
+
responses: list[DecisionResponse] = []
|
|
182
|
+
offset = 0
|
|
183
|
+
for request, request_tasks in zip(requests, compiled, strict=True):
|
|
184
|
+
end = offset + len(request_tasks)
|
|
185
|
+
responses.append(DecisionResponse(
|
|
186
|
+
model=self.model_name,
|
|
187
|
+
answers=aggregate_margin_answers(request, request_tasks, margins[offset:end], self.metadata),
|
|
188
|
+
usage=Usage(input_tokens=sum(token_counts[offset:end])),
|
|
189
|
+
))
|
|
190
|
+
offset = end
|
|
191
|
+
return responses
|
|
192
|
+
|
|
193
|
+
|
|
194
|
+
@dataclass(frozen=True)
|
|
195
|
+
class RLCDTemperatureCalibrator:
|
|
196
|
+
"""Safe JSON-only temperature calibrator emitted by OpenJev Verdict."""
|
|
197
|
+
|
|
198
|
+
temperature: float
|
|
199
|
+
per_k: Mapping[str, float]
|
|
200
|
+
|
|
201
|
+
def temperature_for(self, candidates: int) -> float:
|
|
202
|
+
value = self.per_k.get(str(candidates), self.temperature)
|
|
203
|
+
if not math.isfinite(value) or value <= 0:
|
|
204
|
+
raise RuntimeErrorBase("Verdict calibrator temperature must be finite and positive")
|
|
205
|
+
return value
|
|
206
|
+
|
|
207
|
+
|
|
208
|
+
def load_rlcd_calibrator(path: str | Path) -> RLCDTemperatureCalibrator:
|
|
209
|
+
"""Load the public ``rlcd-calibrator-v1`` JSON artifact, never pickle data."""
|
|
210
|
+
|
|
211
|
+
try:
|
|
212
|
+
payload = json.loads(Path(path).read_text())
|
|
213
|
+
except (OSError, json.JSONDecodeError) as exc:
|
|
214
|
+
raise RuntimeErrorBase("unable to read Verdict calibrator JSON") from exc
|
|
215
|
+
if not isinstance(payload, dict) or payload.get("format_version") != "rlcd-calibrator-v1":
|
|
216
|
+
raise RuntimeErrorBase("unsupported Verdict calibrator format; expected rlcd-calibrator-v1 JSON")
|
|
217
|
+
temperature = payload.get("temperature")
|
|
218
|
+
raw_per_k = payload.get("per_k", {})
|
|
219
|
+
if not isinstance(temperature, int | float) or isinstance(temperature, bool):
|
|
220
|
+
raise RuntimeErrorBase("Verdict calibrator temperature must be numeric")
|
|
221
|
+
if not isinstance(raw_per_k, dict):
|
|
222
|
+
raise RuntimeErrorBase("Verdict calibrator per_k must be an object")
|
|
223
|
+
per_k: dict[str, float] = {}
|
|
224
|
+
for count, value in raw_per_k.items():
|
|
225
|
+
if not isinstance(count, str) or not isinstance(value, int | float) or isinstance(value, bool):
|
|
226
|
+
raise RuntimeErrorBase("Verdict calibrator per_k must map strings to numbers")
|
|
227
|
+
try:
|
|
228
|
+
cardinality = int(count)
|
|
229
|
+
except ValueError as exc:
|
|
230
|
+
raise RuntimeErrorBase("Verdict calibrator per_k keys must be positive integers") from exc
|
|
231
|
+
if cardinality <= 0 or str(cardinality) != count:
|
|
232
|
+
raise RuntimeErrorBase("Verdict calibrator per_k keys must be positive integers")
|
|
233
|
+
per_k[count] = float(value)
|
|
234
|
+
calibrator = RLCDTemperatureCalibrator(float(temperature), per_k)
|
|
235
|
+
calibrator.temperature_for(1)
|
|
236
|
+
for count in per_k:
|
|
237
|
+
calibrator.temperature_for(int(count))
|
|
238
|
+
return calibrator
|
|
239
|
+
|
|
240
|
+
|
|
241
|
+
def render_gliclass_prompt(question: str, context: str, labels: Sequence[str]) -> str:
|
|
242
|
+
"""Render Verdict's GLiClass marker contract in one model input string."""
|
|
243
|
+
|
|
244
|
+
if not labels:
|
|
245
|
+
raise RuntimeErrorBase("GLiClass requires at least one candidate label")
|
|
246
|
+
return "".join(f"<<LABEL>>{label}" for label in labels) + f"<<SEP>>Question: {question}\n\nContext:\n{context}"
|
|
247
|
+
|
|
248
|
+
|
|
249
|
+
def compile_verdict_tasks(request: DecisionRequest) -> list[MarginTask]:
|
|
250
|
+
"""Compile Jev fields into Verdict's labels plus its explicit abstention route.
|
|
251
|
+
|
|
252
|
+
Verdict's published GLiClass contract has 25 slots: at most 24 substantive
|
|
253
|
+
labels and one final ``insufficient evidence`` label. This intentionally
|
|
254
|
+
does not reuse generic candidate templates because Verdict's label wording
|
|
255
|
+
is part of the checkpoint's input contract.
|
|
256
|
+
"""
|
|
257
|
+
|
|
258
|
+
query = render_content(request.state)
|
|
259
|
+
tasks: list[MarginTask] = []
|
|
260
|
+
for question_id, question in request.questions.items():
|
|
261
|
+
instruction = render_content(question.instructions)
|
|
262
|
+
substantive: list[tuple[str, str]]
|
|
263
|
+
if isinstance(question, ChoiceQuestion):
|
|
264
|
+
substantive = [
|
|
265
|
+
(option_id, f"It is {render_content(criterion)}")
|
|
266
|
+
for option_id, criterion in question.criteria.items()
|
|
267
|
+
]
|
|
268
|
+
elif isinstance(question, ScoreQuestion):
|
|
269
|
+
substantive = [
|
|
270
|
+
(str(index), f"{render_content(criterion)} (Value: {index})")
|
|
271
|
+
for index, criterion in enumerate(question.criteria)
|
|
272
|
+
]
|
|
273
|
+
elif isinstance(question, NoulQuestion):
|
|
274
|
+
if question.criteria is None:
|
|
275
|
+
proposition = instruction
|
|
276
|
+
substantive = [
|
|
277
|
+
("true", f"true: {proposition}"),
|
|
278
|
+
("false", f"false: not {proposition}"),
|
|
279
|
+
]
|
|
280
|
+
else:
|
|
281
|
+
substantive = [
|
|
282
|
+
("true", f"true: {render_content(question.criteria.true)}"),
|
|
283
|
+
("false", f"false: not {render_content(question.criteria.false)}"),
|
|
284
|
+
]
|
|
285
|
+
else: # pragma: no cover - exhaustive protocol union
|
|
286
|
+
raise RuntimeErrorBase(f"unsupported question type: {type(question).__name__}")
|
|
287
|
+
if len(substantive) > VERDICT_MAX_SUBSTANTIVE_CANDIDATES:
|
|
288
|
+
raise RuntimeErrorBase(
|
|
289
|
+
"Verdict supports at most 24 substantive candidates plus abstention"
|
|
290
|
+
)
|
|
291
|
+
tasks.extend(
|
|
292
|
+
MarginTask(question_id, option_id, instruction, query, document)
|
|
293
|
+
for option_id, document in substantive
|
|
294
|
+
)
|
|
295
|
+
tasks.append(
|
|
296
|
+
MarginTask(
|
|
297
|
+
question_id,
|
|
298
|
+
VERDICT_ABSTENTION_ID,
|
|
299
|
+
instruction,
|
|
300
|
+
query,
|
|
301
|
+
VERDICT_ABSTENTION_LABEL,
|
|
302
|
+
)
|
|
303
|
+
)
|
|
304
|
+
return tasks
|
|
305
|
+
|
|
306
|
+
|
|
307
|
+
def aggregate_verdict_answers(
|
|
308
|
+
request: DecisionRequest,
|
|
309
|
+
tasks: Sequence[MarginTask],
|
|
310
|
+
margins: Sequence[float],
|
|
311
|
+
) -> dict[str, Any]:
|
|
312
|
+
"""Map Verdict's calibrated full distributions to the Jev response types.
|
|
313
|
+
|
|
314
|
+
Choice and score responses retain the explicit abstention probability under
|
|
315
|
+
``__insufficient_evidence__``. Jev's current Noul wire type has no field
|
|
316
|
+
for that mass, so it exposes Verdict's published conditional probability
|
|
317
|
+
``P(true | sufficient evidence)`` instead.
|
|
318
|
+
"""
|
|
319
|
+
|
|
320
|
+
if len(tasks) != len(margins) or not all(math.isfinite(value) for value in margins):
|
|
321
|
+
raise RuntimeErrorBase("Verdict returned invalid distribution logits")
|
|
322
|
+
grouped: dict[str, list[tuple[str, float]]] = {
|
|
323
|
+
question_id: [] for question_id in request.questions
|
|
324
|
+
}
|
|
325
|
+
for task, margin in zip(tasks, margins, strict=True):
|
|
326
|
+
grouped[task.question_id].append((task.option_id, margin))
|
|
327
|
+
|
|
328
|
+
answers: dict[str, Any] = {}
|
|
329
|
+
for question_id, question in request.questions.items():
|
|
330
|
+
candidates = grouped[question_id]
|
|
331
|
+
option_ids = [option_id for option_id, _ in candidates]
|
|
332
|
+
if option_ids.count(VERDICT_ABSTENTION_ID) != 1:
|
|
333
|
+
raise RuntimeErrorBase("Verdict field must include exactly one abstention candidate")
|
|
334
|
+
probabilities = softmax([margin for _, margin in candidates])
|
|
335
|
+
distribution = dict(zip(option_ids, probabilities, strict=True))
|
|
336
|
+
confidence = normalized_entropy_confidence(probabilities)
|
|
337
|
+
if isinstance(question, ChoiceQuestion):
|
|
338
|
+
answers[question_id] = ChoiceAnswer(
|
|
339
|
+
type="choice",
|
|
340
|
+
choice=max(distribution, key=distribution.__getitem__),
|
|
341
|
+
probabilities=distribution,
|
|
342
|
+
confidence=confidence,
|
|
343
|
+
)
|
|
344
|
+
elif isinstance(question, ScoreQuestion):
|
|
345
|
+
substantive = probabilities[:-1]
|
|
346
|
+
substantive_mass = math.fsum(substantive)
|
|
347
|
+
if substantive_mass <= 0:
|
|
348
|
+
raise RuntimeErrorBase("Verdict score distribution has no substantive mass")
|
|
349
|
+
answers[question_id] = ScoreAnswer(
|
|
350
|
+
type="score",
|
|
351
|
+
score=math.fsum(
|
|
352
|
+
index * (probability / substantive_mass)
|
|
353
|
+
for index, probability in enumerate(substantive)
|
|
354
|
+
),
|
|
355
|
+
probabilities=distribution,
|
|
356
|
+
confidence=confidence,
|
|
357
|
+
legend=question.criteria,
|
|
358
|
+
)
|
|
359
|
+
elif isinstance(question, NoulQuestion):
|
|
360
|
+
by_id = distribution
|
|
361
|
+
true_false_mass = by_id["true"] + by_id["false"]
|
|
362
|
+
if true_false_mass <= 0:
|
|
363
|
+
raise RuntimeErrorBase("Verdict noul distribution has no true/false mass")
|
|
364
|
+
answers[question_id] = NoulAnswer(
|
|
365
|
+
type="noul", noul=by_id["true"] / true_false_mass
|
|
366
|
+
)
|
|
367
|
+
else: # pragma: no cover - exhaustive protocol union
|
|
368
|
+
raise RuntimeErrorBase(f"unsupported question type: {type(question).__name__}")
|
|
369
|
+
return answers
|
|
370
|
+
|
|
371
|
+
|
|
372
|
+
class GLiClassCalibratedBackend(DecisionRuntime):
|
|
373
|
+
"""One-pass GLiClass distribution head with optional RLCD temperature scaling."""
|
|
374
|
+
|
|
375
|
+
def __init__(
|
|
376
|
+
self, model_id: str, *, config: dict[str, Any] | None = None, device: str = "auto"
|
|
377
|
+
) -> None:
|
|
378
|
+
try:
|
|
379
|
+
import torch
|
|
380
|
+
from gliclass import GLiClassModel
|
|
381
|
+
from huggingface_hub import hf_hub_download
|
|
382
|
+
from transformers import AutoTokenizer
|
|
383
|
+
except ImportError as exc: # pragma: no cover - optional dependency
|
|
384
|
+
raise RuntimeErrorBase(
|
|
385
|
+
"GLiClassCalibratedBackend requires transformers, torch, huggingface-hub, and gliclass"
|
|
386
|
+
) from exc
|
|
387
|
+
|
|
388
|
+
self.model_name = str((config or {}).get("model", model_id))
|
|
389
|
+
self.config = config or {}
|
|
390
|
+
self.metadata = decision_metadata(self.config)
|
|
391
|
+
if self.metadata.get("readout") != "gliclass_calibrated":
|
|
392
|
+
raise RuntimeErrorBase("GLiClassCalibratedBackend requires decision.readout=gliclass_calibrated")
|
|
393
|
+
self._torch = torch
|
|
394
|
+
loader = _mapping(self.metadata.get("loader", {}), "decision.loader")
|
|
395
|
+
base_model = loader.get("base_model", model_id)
|
|
396
|
+
tokenizer_id = loader.get("tokenizer", base_model)
|
|
397
|
+
if not isinstance(base_model, str) or not isinstance(tokenizer_id, str):
|
|
398
|
+
raise RuntimeErrorBase("decision.loader base_model and tokenizer must be strings")
|
|
399
|
+
kwargs = _loader_kwargs(loader)
|
|
400
|
+
self._tokenizer = AutoTokenizer.from_pretrained(tokenizer_id, **kwargs)
|
|
401
|
+
self._model = GLiClassModel.from_pretrained(base_model, **kwargs)
|
|
402
|
+
self._model.to(_device(torch, self.metadata.get("device"), device))
|
|
403
|
+
self._model.eval()
|
|
404
|
+
self._input_device = next(self._model.parameters()).device
|
|
405
|
+
input_config = _mapping(self.metadata.get("input"), "decision.input")
|
|
406
|
+
self._max_length = _positive_int(input_config.get("max_length"), "decision.input.max_length", 1024)
|
|
407
|
+
self._max_candidates = _positive_int(input_config.get("max_candidates"), "decision.input.max_candidates", 25)
|
|
408
|
+
calibrator_config = self.metadata.get("calibrator")
|
|
409
|
+
self._calibrator: RLCDTemperatureCalibrator | None = None
|
|
410
|
+
if calibrator_config is not None:
|
|
411
|
+
calibrator_spec = _mapping(calibrator_config, "decision.calibrator")
|
|
412
|
+
filename = calibrator_spec.get("file", "calibrator.json")
|
|
413
|
+
repo = calibrator_spec.get("repo", base_model)
|
|
414
|
+
if not isinstance(filename, str) or not isinstance(repo, str):
|
|
415
|
+
raise RuntimeErrorBase("decision.calibrator file and repo must be strings")
|
|
416
|
+
revision = calibrator_spec.get("revision", kwargs.get("revision"))
|
|
417
|
+
if revision is not None and not isinstance(revision, str):
|
|
418
|
+
raise RuntimeErrorBase("decision.calibrator.revision must be a string")
|
|
419
|
+
local_file = hf_hub_download(repo_id=repo, filename=filename, revision=revision)
|
|
420
|
+
self._calibrator = load_rlcd_calibrator(local_file)
|
|
421
|
+
|
|
422
|
+
def decide_batch(self, requests: Sequence[DecisionRequest]) -> list[DecisionResponse]:
|
|
423
|
+
questions: list[tuple[DecisionRequest, str, Any, list[MarginTask]]] = []
|
|
424
|
+
for request in requests:
|
|
425
|
+
by_question: dict[str, list[MarginTask]] = {}
|
|
426
|
+
for task in compile_verdict_tasks(request):
|
|
427
|
+
by_question.setdefault(task.question_id, []).append(task)
|
|
428
|
+
for question_id, question in request.questions.items():
|
|
429
|
+
tasks = by_question[question_id]
|
|
430
|
+
if len(tasks) > self._max_candidates:
|
|
431
|
+
raise RuntimeErrorBase(
|
|
432
|
+
f"GLiClass candidate count {len(tasks)} exceeds configured maximum {self._max_candidates}"
|
|
433
|
+
)
|
|
434
|
+
questions.append((request, question_id, question, tasks))
|
|
435
|
+
|
|
436
|
+
prompts = [
|
|
437
|
+
render_gliclass_prompt(tasks[0].instruction, tasks[0].query, [task.document for task in tasks])
|
|
438
|
+
for _, _, _, tasks in questions
|
|
439
|
+
]
|
|
440
|
+
encoded = self._tokenizer(prompts, padding=True, truncation=True, max_length=self._max_length, return_tensors="pt")
|
|
441
|
+
encoded = {name: value.to(self._input_device) for name, value in encoded.items()}
|
|
442
|
+
with self._torch.inference_mode():
|
|
443
|
+
output = self._model(**encoded)
|
|
444
|
+
logits = output.logits.float()
|
|
445
|
+
if logits.ndim != 2 or logits.shape[0] != len(questions):
|
|
446
|
+
raise RuntimeErrorBase("GLiClass model returned an invalid distribution-head shape")
|
|
447
|
+
|
|
448
|
+
answers_by_request: list[dict[str, Any]] = [{} for _ in requests]
|
|
449
|
+
token_counts = [int(value) for value in encoded["attention_mask"].sum(dim=1).cpu().tolist()]
|
|
450
|
+
request_indices = {id(request): index for index, request in enumerate(requests)}
|
|
451
|
+
for row, (request, question_id, question, tasks) in enumerate(questions):
|
|
452
|
+
if logits.shape[1] < len(tasks):
|
|
453
|
+
raise RuntimeErrorBase("GLiClass distribution head has fewer logits than candidates")
|
|
454
|
+
margins = logits[row, : len(tasks)]
|
|
455
|
+
if not bool(margins.isfinite().all()):
|
|
456
|
+
raise RuntimeErrorBase("GLiClass model produced non-finite candidate logits")
|
|
457
|
+
if self._calibrator is not None:
|
|
458
|
+
margins = margins / self._calibrator.temperature_for(len(tasks))
|
|
459
|
+
single_request = request.model_copy(
|
|
460
|
+
update={"questions": {question_id: question}}
|
|
461
|
+
)
|
|
462
|
+
answers_by_request[request_indices[id(request)]].update(
|
|
463
|
+
aggregate_verdict_answers(single_request, tasks, margins.cpu().tolist())
|
|
464
|
+
)
|
|
465
|
+
|
|
466
|
+
usage_by_request = [0] * len(requests)
|
|
467
|
+
for tokens, (request, _, _, _) in zip(token_counts, questions, strict=True):
|
|
468
|
+
usage_by_request[request_indices[id(request)]] += tokens
|
|
469
|
+
return [
|
|
470
|
+
DecisionResponse(model=self.model_name, answers=answers_by_request[index], usage=Usage(input_tokens=usage_by_request[index]))
|
|
471
|
+
for index in range(len(requests))
|
|
472
|
+
]
|
|
@@ -0,0 +1,71 @@
|
|
|
1
|
+
"""Sentence-Transformers CrossEncoder adapter for open rerankers."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from collections.abc import Sequence
|
|
6
|
+
from typing import Any
|
|
7
|
+
|
|
8
|
+
from .encoder_decoder import (
|
|
9
|
+
aggregate_margin_answers,
|
|
10
|
+
compile_margin_tasks,
|
|
11
|
+
decision_metadata,
|
|
12
|
+
)
|
|
13
|
+
from .protocol import DecisionRequest, DecisionResponse, Usage
|
|
14
|
+
from .runtime import DecisionRuntime, RuntimeErrorBase
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class CrossEncoderBackend(DecisionRuntime):
|
|
18
|
+
"""Expose a public CrossEncoder/reranker as typed candidate probabilities."""
|
|
19
|
+
|
|
20
|
+
def __init__(
|
|
21
|
+
self,
|
|
22
|
+
model_id: str,
|
|
23
|
+
*,
|
|
24
|
+
config: dict[str, Any] | None = None,
|
|
25
|
+
device: str = "auto",
|
|
26
|
+
) -> None:
|
|
27
|
+
try:
|
|
28
|
+
import torch
|
|
29
|
+
from sentence_transformers import CrossEncoder
|
|
30
|
+
except ImportError as exc: # pragma: no cover - optional dependency
|
|
31
|
+
raise RuntimeErrorBase(
|
|
32
|
+
"CrossEncoderBackend requires sentence-transformers and torch"
|
|
33
|
+
) from exc
|
|
34
|
+
|
|
35
|
+
self.model_name = str((config or {}).get("model", model_id))
|
|
36
|
+
self.config = config or {}
|
|
37
|
+
metadata = decision_metadata(self.config)
|
|
38
|
+
if metadata.get("readout") != "cross_encoder_margin":
|
|
39
|
+
raise RuntimeErrorBase(
|
|
40
|
+
"CrossEncoderBackend requires decision.readout=cross_encoder_margin"
|
|
41
|
+
)
|
|
42
|
+
self._batch_size = int(metadata.get("batch_size", 32))
|
|
43
|
+
target = "cuda" if device == "auto" and torch.cuda.is_available() else "cpu"
|
|
44
|
+
self._model = CrossEncoder(
|
|
45
|
+
model_id,
|
|
46
|
+
max_length=int(metadata.get("max_length", 512)),
|
|
47
|
+
device=target,
|
|
48
|
+
)
|
|
49
|
+
|
|
50
|
+
def decide_batch(self, requests: Sequence[DecisionRequest]) -> list[DecisionResponse]:
|
|
51
|
+
compiled = [compile_margin_tasks(request, decision_metadata(self.config)) for request in requests]
|
|
52
|
+
tasks = [task for request_tasks in compiled for task in request_tasks]
|
|
53
|
+
pairs = [(task.query, f"{task.instruction}\n\n{task.document}") for task in tasks]
|
|
54
|
+
scores = self._model.predict(pairs, batch_size=self._batch_size, show_progress_bar=False)
|
|
55
|
+
margins = [float(value) for value in scores]
|
|
56
|
+
responses: list[DecisionResponse] = []
|
|
57
|
+
offset = 0
|
|
58
|
+
metadata = decision_metadata(self.config)
|
|
59
|
+
for request, request_tasks in zip(requests, compiled, strict=True):
|
|
60
|
+
end = offset + len(request_tasks)
|
|
61
|
+
responses.append(
|
|
62
|
+
DecisionResponse(
|
|
63
|
+
model=self.model_name,
|
|
64
|
+
answers=aggregate_margin_answers(
|
|
65
|
+
request, request_tasks, margins[offset:end], metadata
|
|
66
|
+
),
|
|
67
|
+
usage=Usage(),
|
|
68
|
+
)
|
|
69
|
+
)
|
|
70
|
+
offset = end
|
|
71
|
+
return responses
|