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,704 @@
1
+ """Configuration-driven LoRA plus calibrated custom-head decision readouts.
2
+
3
+ The published Open-Jev and SmallJev checkpoints are not ordinary generation
4
+ checkpoints: each combines a pinned base model, a PEFT adapter, a separately
5
+ saved head, and a calibration artifact. This module owns that composition
6
+ without importing an author's serving package or branching on a model name.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import json
12
+ import math
13
+ from collections.abc import Mapping, Sequence
14
+ from dataclasses import dataclass
15
+ from pathlib import Path
16
+ from typing import Any
17
+
18
+ from .encoder_decoder import (
19
+ _mapping,
20
+ _template,
21
+ aggregate_margin_answers,
22
+ compile_margin_tasks,
23
+ decision_metadata,
24
+ )
25
+ from .hidden_state_probe import render_probe_task
26
+ from .protocol import (
27
+ ChoiceAnswer,
28
+ ChoiceQuestion,
29
+ DecisionRequest,
30
+ DecisionResponse,
31
+ NoulAnswer,
32
+ NoulQuestion,
33
+ ScoreAnswer,
34
+ ScoreQuestion,
35
+ Usage,
36
+ )
37
+ from .runtime import DecisionRuntime, RuntimeErrorBase, softmax
38
+
39
+
40
+ def _positive_number(value: Any, name: str) -> float:
41
+ if isinstance(value, bool) or not isinstance(value, int | float):
42
+ raise RuntimeErrorBase(f"{name} must be a positive finite number")
43
+ result = float(value)
44
+ if not math.isfinite(result) or result <= 0.0:
45
+ raise RuntimeErrorBase(f"{name} must be a positive finite number")
46
+ return result
47
+
48
+
49
+ def calibration_temperature(payload: Mapping[str, Any], field: str) -> float:
50
+ """Extract the one saved temperature that calibrates every candidate logit."""
51
+
52
+ if not isinstance(field, str) or not field:
53
+ raise RuntimeErrorBase("decision.artifacts.calibration.field must be a string")
54
+ if field not in payload:
55
+ raise RuntimeErrorBase(f"calibration artifact is missing field: {field}")
56
+ return _positive_number(payload[field], f"calibration artifact field {field!r}")
57
+
58
+
59
+ def custom_head_metadata(config: Mapping[str, Any]) -> dict[str, Any]:
60
+ metadata = decision_metadata(config)
61
+ if metadata.get("readout") not in {"openjev_scalar_head", "semantic_option_head"}:
62
+ raise RuntimeErrorBase(
63
+ "custom-head backend requires decision.readout=openjev_scalar_head "
64
+ "or semantic_option_head"
65
+ )
66
+ _mapping(metadata.get("loader"), "decision.loader")
67
+ _mapping(metadata.get("artifacts"), "decision.artifacts")
68
+ _mapping(metadata.get("input"), "decision.input")
69
+ return metadata
70
+
71
+
72
+ def render_custom_head_task(task: Any, metadata: Mapping[str, Any]) -> str:
73
+ """Render a candidate exactly once from service-owned input metadata."""
74
+
75
+ return render_probe_task(task, metadata)
76
+
77
+
78
+ def _openjev_content(value: Any) -> str:
79
+ if isinstance(value, str):
80
+ return value
81
+ try:
82
+ return json.dumps(value, ensure_ascii=False, sort_keys=True, allow_nan=False)
83
+ except (TypeError, ValueError) as exc:
84
+ raise RuntimeErrorBase("Open-Jev content must be JSON serializable") from exc
85
+
86
+
87
+ @dataclass(frozen=True)
88
+ class OpenJevTask:
89
+ question_id: str
90
+ question: ChoiceQuestion | ScoreQuestion | NoulQuestion
91
+ labels: tuple[str, ...]
92
+ prompts: tuple[str, ...]
93
+
94
+
95
+ def compile_openjev_tasks(request: DecisionRequest) -> list[OpenJevTask]:
96
+ """Match Open-Jev's isolated candidate and single-Noul prompt contract."""
97
+
98
+ state = _openjev_content(request.state)
99
+ tasks: list[OpenJevTask] = []
100
+ for name, question in request.questions.items():
101
+ instruction = _openjev_content(question.instructions)
102
+ if isinstance(question, NoulQuestion) and question.criteria is not None:
103
+ instruction += (
104
+ f"\nYes means: {_openjev_content(question.criteria.true)}"
105
+ f"\nNo means: {_openjev_content(question.criteria.false)}"
106
+ )
107
+ prefix = f"Context:\n{state}\n\nQuestion: {instruction}\n"
108
+ if isinstance(question, NoulQuestion):
109
+ tasks.append(
110
+ OpenJevTask(
111
+ name,
112
+ question,
113
+ ("false", "true"),
114
+ (prefix + "Is the answer to this question yes? Answer Yes or No.",),
115
+ )
116
+ )
117
+ continue
118
+ if isinstance(question, ChoiceQuestion):
119
+ labels = tuple(question.criteria)
120
+ options = tuple(
121
+ key if value is None else f"{key}: {_openjev_content(value)}"
122
+ for key, value in question.criteria.items()
123
+ )
124
+ else:
125
+ if len(question.criteria) > 10:
126
+ raise RuntimeErrorBase("Open-Jev Score supports at most 10 levels")
127
+ labels = tuple(str(index) for index in range(len(question.criteria)))
128
+ options = tuple(_openjev_content(value) for value in question.criteria)
129
+ prompts = tuple(
130
+ prefix
131
+ + f"Proposed answer: {option}\n"
132
+ + "Is this proposed answer correct? Answer Yes or No."
133
+ for option in options
134
+ )
135
+ tasks.append(OpenJevTask(name, question, labels, prompts))
136
+ return tasks
137
+
138
+
139
+ def format_openjev_answers(
140
+ tasks: Sequence[OpenJevTask], scores: Sequence[float], temperature: float
141
+ ) -> dict[str, Any]:
142
+ """Apply the checkpoint temperature before each typed Open-Jev readout."""
143
+
144
+ temperature = _positive_number(temperature, "Open-Jev temperature")
145
+ answers: dict[str, Any] = {}
146
+ offset = 0
147
+ for task in tasks:
148
+ size = len(task.prompts)
149
+ values = scores[offset : offset + size]
150
+ if len(values) != size or not all(math.isfinite(value) for value in values):
151
+ raise RuntimeErrorBase("Open-Jev scorer returned invalid candidate logits")
152
+ offset += size
153
+ if isinstance(task.question, NoulQuestion):
154
+ probability = softmax([0.0, values[0] / temperature])[1]
155
+ answers[task.question_id] = NoulAnswer(type="noul", noul=probability)
156
+ continue
157
+ probabilities = softmax([value / temperature for value in values])
158
+ distribution = dict(zip(task.labels, probabilities, strict=True))
159
+ confidence = max(probabilities)
160
+ if isinstance(task.question, ChoiceQuestion):
161
+ answers[task.question_id] = ChoiceAnswer(
162
+ type="choice",
163
+ choice=max(distribution, key=distribution.__getitem__),
164
+ probabilities=distribution,
165
+ confidence=confidence,
166
+ )
167
+ else:
168
+ answers[task.question_id] = ScoreAnswer(
169
+ type="score",
170
+ score=math.fsum(index * probability for index, probability in enumerate(probabilities)),
171
+ probabilities=distribution,
172
+ confidence=confidence,
173
+ legend=task.question.criteria,
174
+ )
175
+ if offset != len(scores):
176
+ raise RuntimeErrorBase("Open-Jev scorer returned extra candidate logits")
177
+ return answers
178
+
179
+
180
+ class ConfiguredCustomHeadBackend(DecisionRuntime):
181
+ """Score independently rendered candidates through a LoRA and saved head.
182
+
183
+ ``decision.head.kind=linear`` implements Open-Jev's saved ``nn.Linear``
184
+ state dict. ``semantic_layers`` evaluates a named sequence of linear
185
+ tensor layers, allowing a SmallJev-style semantic OptionScorerHead to be
186
+ expressed by checkpoint keys rather than a copied upstream Python class.
187
+ """
188
+
189
+ def __init__(
190
+ self,
191
+ model_id: str,
192
+ *,
193
+ config: dict[str, Any] | None = None,
194
+ device: str = "auto",
195
+ ) -> None:
196
+ try:
197
+ import torch
198
+ from huggingface_hub import hf_hub_download
199
+ from peft import PeftModel
200
+ from transformers import (
201
+ AutoModel,
202
+ AutoModelForImageTextToText,
203
+ AutoTokenizer,
204
+ )
205
+ except ImportError as exc: # pragma: no cover - optional dependency
206
+ raise RuntimeErrorBase(
207
+ "custom-head backends require transformers, torch, peft, and "
208
+ "huggingface-hub"
209
+ ) from exc
210
+
211
+ self.model_name = str((config or {}).get("model", model_id))
212
+ self.config = config or {}
213
+ self.metadata = custom_head_metadata(self.config)
214
+ self._torch = torch
215
+ loader = _mapping(self.metadata["loader"], "decision.loader")
216
+ base_model = loader.get("base_model", model_id)
217
+ if not isinstance(base_model, str):
218
+ raise RuntimeErrorBase("decision.loader.base_model must be a string")
219
+ revision = loader.get("revision")
220
+ if revision is not None and not isinstance(revision, str):
221
+ raise RuntimeErrorBase("decision.loader.revision must be a string")
222
+ trust_remote_code = loader.get("trust_remote_code", False)
223
+ if not isinstance(trust_remote_code, bool):
224
+ raise RuntimeErrorBase("decision.loader.trust_remote_code must be boolean")
225
+ tokenizer_id = loader.get("tokenizer", base_model)
226
+ if not isinstance(tokenizer_id, str):
227
+ raise RuntimeErrorBase("decision.loader.tokenizer must be a string")
228
+ requested_device = self.metadata.get("device", device)
229
+ target = self._resolve_device(requested_device)
230
+ dtype = self._resolve_dtype(self.metadata.get("dtype"), target)
231
+ common: dict[str, Any] = {"trust_remote_code": trust_remote_code}
232
+ if revision is not None:
233
+ common["revision"] = revision
234
+ self._tokenizer = AutoTokenizer.from_pretrained(tokenizer_id, **common)
235
+ if self._tokenizer.pad_token_id is None:
236
+ if self._tokenizer.eos_token_id is None:
237
+ raise RuntimeErrorBase("tokenizer must define a pad token or an EOS token")
238
+ self._tokenizer.pad_token = self._tokenizer.eos_token
239
+ self._tokenizer.padding_side = "right"
240
+
241
+ model_kwargs = dict(common)
242
+ model_kwargs["dtype"] = dtype
243
+ attention = loader.get("attn_implementation", "sdpa")
244
+ if attention is not None:
245
+ if not isinstance(attention, str):
246
+ raise RuntimeErrorBase("decision.loader.attn_implementation must be a string or null")
247
+ model_kwargs["attn_implementation"] = attention
248
+ model_class = loader.get("model_class", "auto")
249
+ if model_class == "auto":
250
+ loaded_model = AutoModel.from_pretrained(base_model, **model_kwargs)
251
+ elif model_class == "image_text_to_text":
252
+ loaded_model = AutoModelForImageTextToText.from_pretrained(
253
+ base_model, **model_kwargs
254
+ )
255
+ else:
256
+ raise RuntimeErrorBase(
257
+ "decision.loader.model_class must be auto or image_text_to_text"
258
+ )
259
+ backbone_path = loader.get("backbone_path")
260
+ self._model = self._resolve_backbone(loaded_model, backbone_path)
261
+ adapter = _mapping(loader.get("adapter"), "decision.loader.adapter")
262
+ adapter_repo = adapter.get("repo")
263
+ if not isinstance(adapter_repo, str):
264
+ raise RuntimeErrorBase("decision.loader.adapter.repo must be a string")
265
+ adapter_kwargs: dict[str, Any] = {}
266
+ adapter_revision = adapter.get("revision")
267
+ if adapter_revision is not None:
268
+ if not isinstance(adapter_revision, str):
269
+ raise RuntimeErrorBase("decision.loader.adapter.revision must be a string")
270
+ adapter_kwargs["revision"] = adapter_revision
271
+ adapter_subfolder = adapter.get("subfolder")
272
+ if adapter_subfolder is not None:
273
+ if not isinstance(adapter_subfolder, str):
274
+ raise RuntimeErrorBase("decision.loader.adapter.subfolder must be a string")
275
+ adapter_kwargs["subfolder"] = adapter_subfolder
276
+ self._model = PeftModel.from_pretrained(self._model, adapter_repo, **adapter_kwargs)
277
+ self._model.to(target)
278
+ self._model.eval()
279
+ self._device = next(self._model.parameters()).device
280
+
281
+ input_config = _mapping(self.metadata["input"], "decision.input")
282
+ _template(input_config.get("template"), "decision.input.template")
283
+ self._max_length = self._positive_int(input_config.get("max_length"), "decision.input.max_length")
284
+ self._batch_size = self._positive_int(self.metadata.get("batch_size", 8), "decision.batch_size")
285
+ chat = _mapping(input_config.get("chat_template", {}), "decision.input.chat_template")
286
+ self._add_generation_prompt = chat.get("add_generation_prompt", True)
287
+ self._enable_thinking = chat.get("enable_thinking", False)
288
+ if not isinstance(self._add_generation_prompt, bool) or not isinstance(self._enable_thinking, bool):
289
+ raise RuntimeErrorBase("decision.input.chat_template flags must be boolean")
290
+
291
+ artifacts = _mapping(self.metadata["artifacts"], "decision.artifacts")
292
+ head_config = _mapping(artifacts.get("head"), "decision.artifacts.head")
293
+ self._head_state = self._download_torch_state(hf_hub_download, head_config, "head")
294
+ self._head = self._build_head(head_config)
295
+ calibration = _mapping(artifacts.get("calibration"), "decision.artifacts.calibration")
296
+ calibration_payload = self._download_json(hf_hub_download, calibration, "calibration")
297
+ field = calibration.get("field", "temperature")
298
+ self._temperature = calibration_temperature(calibration_payload, field)
299
+
300
+ @staticmethod
301
+ def _positive_int(value: Any, name: str) -> int:
302
+ if not isinstance(value, int) or isinstance(value, bool) or value <= 0:
303
+ raise RuntimeErrorBase(f"{name} must be a positive integer")
304
+ return value
305
+
306
+ @staticmethod
307
+ def _resolve_backbone(model: Any, path: Any) -> Any:
308
+ if path is None:
309
+ return model
310
+ if not isinstance(path, str) or not path:
311
+ raise RuntimeErrorBase("decision.loader.backbone_path must be a non-empty string")
312
+ value = model
313
+ for part in path.split("."):
314
+ if not hasattr(value, part):
315
+ raise RuntimeErrorBase(
316
+ f"decision.loader.backbone_path is missing component: {part}"
317
+ )
318
+ value = getattr(value, part)
319
+ return value
320
+
321
+ def _resolve_device(self, requested: Any) -> str:
322
+ torch = self._torch
323
+ target = "cuda" if requested == "auto" and torch.cuda.is_available() else "cpu" if requested == "auto" else requested
324
+ if not isinstance(target, str):
325
+ raise RuntimeErrorBase("decision.device must be a string")
326
+ if target.startswith("cuda") and not torch.cuda.is_available():
327
+ raise RuntimeErrorBase("CUDA was requested, but no CUDA device is available")
328
+ return target
329
+
330
+ def _resolve_dtype(self, value: Any, device: str) -> Any:
331
+ torch = self._torch
332
+ if value is None or value == "auto":
333
+ return torch.bfloat16 if device.startswith("cuda") else torch.float32
334
+ mapping = {"bfloat16": torch.bfloat16, "bf16": torch.bfloat16, "float16": torch.float16, "fp16": torch.float16, "float32": torch.float32, "fp32": torch.float32}
335
+ if value not in mapping:
336
+ raise RuntimeErrorBase(f"unsupported decision dtype: {value!r}")
337
+ return mapping[value]
338
+
339
+ @staticmethod
340
+ def _artifact_download_kwargs(config: Mapping[str, Any], name: str) -> dict[str, str]:
341
+ repo = config.get("repo")
342
+ file = config.get("file")
343
+ if not isinstance(repo, str) or not isinstance(file, str):
344
+ raise RuntimeErrorBase(f"decision.artifacts.{name} requires string repo and file")
345
+ result = {"repo_id": repo, "filename": file}
346
+ revision = config.get("revision")
347
+ if revision is not None:
348
+ if not isinstance(revision, str):
349
+ raise RuntimeErrorBase(f"decision.artifacts.{name}.revision must be a string")
350
+ result["revision"] = revision
351
+ return result
352
+
353
+ def _download_torch_state(self, download: Any, config: Mapping[str, Any], name: str) -> Mapping[str, Any]:
354
+ path = download(**self._artifact_download_kwargs(config, name))
355
+ payload = self._torch.load(path, map_location="cpu", weights_only=True)
356
+ state_key = config.get("state_key")
357
+ if state_key is not None:
358
+ if not isinstance(state_key, str) or not isinstance(payload, Mapping) or state_key not in payload:
359
+ raise RuntimeErrorBase(f"decision.artifacts.{name}.state_key is missing from artifact")
360
+ payload = payload[state_key]
361
+ if not isinstance(payload, Mapping):
362
+ raise RuntimeErrorBase(f"decision {name} artifact must be a tensor state mapping")
363
+ return payload
364
+
365
+ def _download_json(self, download: Any, config: Mapping[str, Any], name: str) -> Mapping[str, Any]:
366
+ path = download(**self._artifact_download_kwargs(config, name))
367
+ try:
368
+ payload = json.loads(Path(path).read_text(encoding="utf-8"))
369
+ except (OSError, json.JSONDecodeError) as exc:
370
+ raise RuntimeErrorBase(f"decision {name} artifact is not valid JSON") from exc
371
+ if not isinstance(payload, Mapping):
372
+ raise RuntimeErrorBase(f"decision {name} artifact must be a JSON object")
373
+ return payload
374
+
375
+ def _build_head(self, config: Mapping[str, Any]) -> Any:
376
+ kind = config.get("kind")
377
+ hidden_size = getattr(self._model.config, "hidden_size", None)
378
+ if not isinstance(hidden_size, int) or hidden_size <= 0:
379
+ raise RuntimeErrorBase("base model config must expose a positive hidden_size")
380
+ torch = self._torch
381
+ if kind == "linear":
382
+ head = torch.nn.Linear(hidden_size, 1, bias=True, dtype=torch.float32)
383
+ try:
384
+ head.load_state_dict(self._head_state, strict=True)
385
+ except (RuntimeError, ValueError) as exc:
386
+ raise RuntimeErrorBase("linear head artifact does not match hidden_size -> 1") from exc
387
+ return head.to(self._device).eval()
388
+ if kind != "semantic_layers":
389
+ raise RuntimeErrorBase("decision.artifacts.head.kind must be linear or semantic_layers")
390
+ layers = config.get("layers")
391
+ if not isinstance(layers, list) or not layers:
392
+ raise RuntimeErrorBase("semantic_layers head requires a non-empty layers array")
393
+ parsed: list[tuple[Any, Any | None, str]] = []
394
+ for index, raw in enumerate(layers):
395
+ layer = _mapping(raw, f"decision.artifacts.head.layers[{index}]")
396
+ weight_key = layer.get("weight")
397
+ bias_key = layer.get("bias")
398
+ activation = layer.get("activation", "identity")
399
+ if not isinstance(weight_key, str) or weight_key not in self._head_state:
400
+ raise RuntimeErrorBase(f"semantic head layer {index} is missing its weight tensor")
401
+ if bias_key is not None and (not isinstance(bias_key, str) or bias_key not in self._head_state):
402
+ raise RuntimeErrorBase(f"semantic head layer {index} is missing its bias tensor")
403
+ if activation not in {"identity", "gelu", "relu", "silu", "tanh"}:
404
+ raise RuntimeErrorBase(f"semantic head layer {index} has unsupported activation")
405
+ parsed.append((self._head_state[weight_key].to(self._device), self._head_state[bias_key].to(self._device) if isinstance(bias_key, str) else None, activation))
406
+ return parsed
407
+
408
+ def _apply_head(self, hidden: Any) -> Any:
409
+ torch = self._torch
410
+ if not isinstance(self._head, list):
411
+ return self._head(hidden.float()).squeeze(-1)
412
+ values = hidden.float()
413
+ for weight, bias, activation in self._head:
414
+ values = torch.nn.functional.linear(values, weight.float(), None if bias is None else bias.float())
415
+ values = {"identity": lambda x: x, "gelu": torch.nn.functional.gelu, "relu": torch.relu, "silu": torch.nn.functional.silu, "tanh": torch.tanh}[activation](values)
416
+ if values.ndim != 1:
417
+ if values.ndim != 2 or values.shape[1] != 1:
418
+ raise RuntimeErrorBase("semantic OptionScorerHead must produce one scalar per candidate")
419
+ values = values[:, 0]
420
+ return values
421
+
422
+ def _score_texts(self, texts: Sequence[str]) -> tuple[list[float], list[int]]:
423
+ torch = self._torch
424
+ scores: list[float] = []
425
+ input_tokens: list[int] = []
426
+ for start in range(0, len(texts), self._batch_size):
427
+ messages = list(texts[start : start + self._batch_size])
428
+ rendered = [self._tokenizer.apply_chat_template([{"role": "user", "content": text}], tokenize=False, add_generation_prompt=self._add_generation_prompt, enable_thinking=self._enable_thinking) for text in messages]
429
+ encoded = self._tokenizer(rendered, return_tensors="pt", padding=True, truncation=False)
430
+ mask = encoded.get("attention_mask")
431
+ if mask is None:
432
+ raise RuntimeErrorBase("tokenizer output has no attention_mask")
433
+ if int(mask.sum(dim=1).max().item()) > self._max_length:
434
+ raise RuntimeErrorBase(f"input length exceeds configured max_length={self._max_length}; no silent truncation")
435
+ input_tokens.extend(int(value) for value in mask.sum(dim=1).tolist())
436
+ encoded = {key: value.to(self._device) for key, value in encoded.items()}
437
+ with torch.inference_mode():
438
+ output = self._model(**encoded, use_cache=False, return_dict=True)
439
+ hidden = getattr(output, "last_hidden_state", None)
440
+ if hidden is None:
441
+ raise RuntimeErrorBase("base model output has no last_hidden_state")
442
+ positions = encoded["attention_mask"].sum(dim=1) - 1
443
+ rows = torch.arange(hidden.shape[0], device=hidden.device)
444
+ values = self._apply_head(hidden[rows, positions])
445
+ if not bool(torch.isfinite(values).all()):
446
+ raise RuntimeErrorBase("custom head produced non-finite candidate logits")
447
+ # Preserve native head logits here. Each typed readout applies the
448
+ # checkpoint's saved temperature exactly once when it constructs
449
+ # the final probability distribution.
450
+ scores.extend(float(value) for value in values.cpu().tolist())
451
+ return scores, input_tokens
452
+
453
+ def decide_batch(self, requests: Sequence[DecisionRequest]) -> list[DecisionResponse]:
454
+ compiled = [compile_margin_tasks(request, self.metadata) for request in requests]
455
+ tasks = [task for request_tasks in compiled for task in request_tasks]
456
+ scores, token_counts = self._score_texts([render_custom_head_task(task, self.metadata) for task in tasks])
457
+ responses: list[DecisionResponse] = []
458
+ offset = 0
459
+ for request, request_tasks in zip(requests, compiled, strict=True):
460
+ end = offset + len(request_tasks)
461
+ responses.append(DecisionResponse(model=self.model_name, answers=aggregate_margin_answers(request, request_tasks, scores[offset:end], self.metadata), usage=Usage(input_tokens=sum(token_counts[offset:end]))))
462
+ offset = end
463
+ return responses
464
+
465
+
466
+ class OpenJevScalarHeadBackend(ConfiguredCustomHeadBackend):
467
+ """Faithful Open-Jev LoRA, scalar-head, and temperature composition."""
468
+
469
+ def decide_batch(self, requests: Sequence[DecisionRequest]) -> list[DecisionResponse]:
470
+ compiled = [compile_openjev_tasks(request) for request in requests]
471
+ tasks = [task for request_tasks in compiled for task in request_tasks]
472
+ prompts = [prompt for task in tasks for prompt in task.prompts]
473
+ scores, token_counts = self._score_texts(prompts)
474
+ responses: list[DecisionResponse] = []
475
+ task_offset = 0
476
+ score_offset = 0
477
+ for request, request_tasks in zip(requests, compiled, strict=True):
478
+ task_end = task_offset + len(request_tasks)
479
+ request_tasks = tasks[task_offset:task_end]
480
+ count = sum(len(task.prompts) for task in request_tasks)
481
+ score_end = score_offset + count
482
+ responses.append(
483
+ DecisionResponse(
484
+ model=self.model_name,
485
+ answers=format_openjev_answers(
486
+ request_tasks, scores[score_offset:score_end], self._temperature
487
+ ),
488
+ usage=Usage(input_tokens=sum(token_counts[score_offset:score_end])),
489
+ )
490
+ )
491
+ task_offset = task_end
492
+ score_offset = score_end
493
+ return responses
494
+
495
+
496
+ def build_smalljev_semantic_ids(
497
+ tokenizer: Any,
498
+ state: str,
499
+ question: str,
500
+ options: Sequence[str],
501
+ *,
502
+ max_length: int = 1024,
503
+ ) -> tuple[list[int], list[tuple[int, int]]]:
504
+ """Exact public SmallJev span construction, including its state-only trim."""
505
+
506
+ if len(options) > 26:
507
+ raise RuntimeErrorBase("SmallJev semantic Choice supports at most 26 options")
508
+ current_state = state
509
+ for _ in range(4):
510
+ head = tokenizer(
511
+ f"State: {current_state}\nQuestion: {question}\nOptions:",
512
+ add_special_tokens=True,
513
+ )["input_ids"]
514
+ chunks = [
515
+ (
516
+ tokenizer(f"\n{chr(ord('A') + index)}.", add_special_tokens=False)[
517
+ "input_ids"
518
+ ],
519
+ tokenizer(f" {option}", add_special_tokens=False)["input_ids"],
520
+ )
521
+ for index, option in enumerate(options)
522
+ ]
523
+ tail = tokenizer("\nAnswer with a single letter:", add_special_tokens=False)[
524
+ "input_ids"
525
+ ]
526
+ total = len(head) + sum(len(marker) + len(text) for marker, text in chunks) + len(tail)
527
+ if total <= max_length or len(current_state) < 100:
528
+ break
529
+ current_state = current_state[: max(50, len(current_state) - int((total - max_length) * 1.5))]
530
+ input_ids = list(head)
531
+ spans: list[tuple[int, int]] = []
532
+ for marker, text in chunks:
533
+ input_ids.extend(marker)
534
+ spans.append((len(input_ids), len(input_ids) + len(text)))
535
+ input_ids.extend(text)
536
+ input_ids.extend(tail)
537
+ return input_ids, spans
538
+
539
+
540
+ class SmallJevSemanticBackend(DecisionRuntime):
541
+ """Faithful published SmallJev semantic-v9 Choice scorer.
542
+
543
+ The public semantic runtime does not apply a saved calibration artifact and
544
+ only uses ``OptionScorerHead`` for Choice. Noul and Score use a separate
545
+ LM-verbalizer path, so this backend deliberately exposes Choice only.
546
+ ``QuestionTypeRuntime`` supplies explicit unsupported responses for the
547
+ remaining wire types when the registry declares that boundary.
548
+ """
549
+
550
+ def __init__(
551
+ self,
552
+ model_id: str,
553
+ *,
554
+ config: dict[str, Any] | None = None,
555
+ device: str = "auto",
556
+ ) -> None:
557
+ try:
558
+ import torch
559
+ from huggingface_hub import hf_hub_download
560
+ from peft import PeftModel
561
+ from transformers import AutoModelForCausalLM, AutoTokenizer
562
+ except ImportError as exc: # pragma: no cover - optional dependency
563
+ raise RuntimeErrorBase(
564
+ "SmallJev semantic backend requires transformers, torch, peft, "
565
+ "and huggingface-hub"
566
+ ) from exc
567
+ self.model_name = str((config or {}).get("model", model_id))
568
+ self.config = config or {}
569
+ self.metadata = decision_metadata(self.config)
570
+ if self.metadata.get("readout") != "semantic_option_head":
571
+ raise RuntimeErrorBase(
572
+ "SmallJevSemanticBackend requires decision.readout=semantic_option_head"
573
+ )
574
+ self._torch = torch
575
+ loader = _mapping(self.metadata.get("loader"), "decision.loader")
576
+ base_model = loader.get("base_model", model_id)
577
+ revision = loader.get("revision")
578
+ if not isinstance(base_model, str) or not isinstance(revision, str):
579
+ raise RuntimeErrorBase(
580
+ "SmallJev loader requires pinned string base_model and revision"
581
+ )
582
+ tokenizer_id = loader.get("tokenizer", base_model)
583
+ if not isinstance(tokenizer_id, str):
584
+ raise RuntimeErrorBase("decision.loader.tokenizer must be a string")
585
+ requested = self.metadata.get("device", device)
586
+ target = "cuda" if requested == "auto" and torch.cuda.is_available() else "cpu" if requested == "auto" else requested
587
+ if not isinstance(target, str) or (target.startswith("cuda") and not torch.cuda.is_available()):
588
+ raise RuntimeErrorBase("requested SmallJev device is unavailable")
589
+ dtype = torch.bfloat16 if target.startswith("cuda") else torch.float32
590
+ self._tokenizer = AutoTokenizer.from_pretrained(tokenizer_id, revision=revision)
591
+ if self._tokenizer.pad_token is None:
592
+ self._tokenizer.pad_token = self._tokenizer.eos_token
593
+ self._model = AutoModelForCausalLM.from_pretrained(
594
+ base_model,
595
+ revision=revision,
596
+ dtype=dtype,
597
+ attn_implementation="eager",
598
+ )
599
+ adapter = _mapping(loader.get("adapter"), "decision.loader.adapter")
600
+ adapter_repo = adapter.get("repo")
601
+ adapter_revision = adapter.get("revision")
602
+ adapter_subfolder = adapter.get("subfolder")
603
+ if not isinstance(adapter_repo, str) or not isinstance(adapter_revision, str) or not isinstance(adapter_subfolder, str):
604
+ raise RuntimeErrorBase(
605
+ "SmallJev adapter requires pinned repo, revision, and subfolder strings"
606
+ )
607
+ self._model = PeftModel.from_pretrained(
608
+ self._model,
609
+ adapter_repo,
610
+ revision=adapter_revision,
611
+ subfolder=adapter_subfolder,
612
+ ).to(target).eval()
613
+ self._device = next(self._model.parameters()).device
614
+ artifacts = _mapping(self.metadata.get("artifacts"), "decision.artifacts")
615
+ head = _mapping(artifacts.get("head"), "decision.artifacts.head")
616
+ repo = head.get("repo")
617
+ file = head.get("file")
618
+ head_revision = head.get("revision")
619
+ if not isinstance(repo, str) or not isinstance(file, str) or not isinstance(head_revision, str):
620
+ raise RuntimeErrorBase("SmallJev head requires pinned repo, file, and revision")
621
+ blob = torch.load(
622
+ hf_hub_download(repo, file, revision=head_revision),
623
+ map_location="cpu",
624
+ weights_only=True,
625
+ )
626
+ if not isinstance(blob, Mapping) or not isinstance(blob.get("hidden_size"), int):
627
+ raise RuntimeErrorBase("SmallJev OptionScorerHead artifact is malformed")
628
+ state = blob.get("state_dict")
629
+ if not isinstance(state, Mapping):
630
+ raise RuntimeErrorBase("SmallJev OptionScorerHead lacks state_dict")
631
+ linear_state = {
632
+ name: state.get(name, state.get(f"scorer.{name}"))
633
+ for name in ("weight", "bias")
634
+ }
635
+ if any(value is None for value in linear_state.values()):
636
+ raise RuntimeErrorBase(
637
+ "SmallJev OptionScorerHead lacks scorer.weight or scorer.bias"
638
+ )
639
+ self._head = torch.nn.Linear(int(blob["hidden_size"]), 1)
640
+ try:
641
+ self._head.load_state_dict(linear_state, strict=True)
642
+ except RuntimeError as exc:
643
+ raise RuntimeErrorBase("SmallJev OptionScorerHead state is incompatible") from exc
644
+ self._head.to(self._device).eval()
645
+ self._max_length = 1024
646
+
647
+ def decide_batch(self, requests: Sequence[DecisionRequest]) -> list[DecisionResponse]:
648
+ torch = self._torch
649
+ responses: list[DecisionResponse] = []
650
+ for request in requests:
651
+ answers: dict[str, Any] = {}
652
+ input_tokens = 0
653
+ state = _openjev_content(request.state)
654
+ for name, question in request.questions.items():
655
+ if not isinstance(question, ChoiceQuestion):
656
+ raise RuntimeErrorBase(
657
+ "SmallJev semantic-v9 faithfully supports Choice only; "
658
+ "declare decision.question_types=[\"choice\"]"
659
+ )
660
+ labels = list(question.criteria)
661
+ options = [
662
+ key if value is None else f"{key}: {_openjev_content(value)}"
663
+ for key, value in question.criteria.items()
664
+ ]
665
+ ids, spans = build_smalljev_semantic_ids(
666
+ self._tokenizer,
667
+ state,
668
+ _openjev_content(question.instructions),
669
+ options,
670
+ max_length=self._max_length,
671
+ )
672
+ input_tokens += len(ids)
673
+ tensor = torch.tensor([ids], device=self._device)
674
+ with torch.inference_mode():
675
+ output = self._model(
676
+ input_ids=tensor,
677
+ use_cache=False,
678
+ output_hidden_states=True,
679
+ )
680
+ hidden = output.hidden_states[-1][0]
681
+ representations = [
682
+ hidden[-1, :] if end <= start else hidden[start:end, :].float().mean(0)
683
+ for start, end in spans
684
+ ]
685
+ with torch.inference_mode():
686
+ logits = self._head(torch.stack(representations).float()).squeeze(-1)
687
+ if not bool(torch.isfinite(logits).all()):
688
+ raise RuntimeErrorBase("SmallJev OptionScorerHead produced non-finite logits")
689
+ probabilities = softmax([float(value) for value in logits.cpu().tolist()])
690
+ distribution = dict(zip(labels, probabilities, strict=True))
691
+ answers[name] = ChoiceAnswer(
692
+ type="choice",
693
+ choice=max(distribution, key=distribution.__getitem__),
694
+ probabilities=distribution,
695
+ confidence=max(probabilities),
696
+ )
697
+ responses.append(
698
+ DecisionResponse(
699
+ model=self.model_name,
700
+ answers=answers,
701
+ usage=Usage(input_tokens=input_tokens),
702
+ )
703
+ )
704
+ return responses