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.
@@ -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