"""Run JADE's decision head with vLLM.""" from __future__ import annotations import hashlib import json import math from pathlib import Path from typing import Any from .decision_prompt import decision_messages, options MAX_OPTIONS = 255 MAX_TOKENS = 8192 LOGPROB_CHUNK_SIZE = 128 class UnsupportedInput(ValueError): """The request exceeds the model's supported types or capacity.""" def distribution(values: list[float], temperature: float) -> list[float]: if not values or not math.isfinite(temperature) or temperature <= 0: raise ValueError("Invalid temperature or missing logits") if any(not math.isfinite(value) for value in values): raise ValueError("Nonfinite answer logits") top = max(values) weights = [math.exp((value - top) / temperature) for value in values] total = sum(weights) return [value / total for value in weights] def validate_release(root: Path) -> dict: manifest = json.loads((root / "release-manifest.json").read_text()) for name, expected in manifest["files"].items(): path = (root / name).resolve() if not path.is_relative_to(root.resolve()): raise ValueError("Release contains a path outside its directory") if not path.is_file(): raise ValueError(f"Missing release file: {name}") with path.open("rb") as stream: actual = hashlib.file_digest(stream, "sha256").hexdigest() if actual != expected: raise ValueError(f"Release file checksum mismatch: {name}") return manifest class JadeEngine: """Load a JADE checkpoint and answer Choice or Noul questions.""" def __init__( self, model: str, revision: str | None = None, max_tokens: int = MAX_TOKENS, gpu_memory_utilization: float = 0.85, ): if not 1 <= int(max_tokens) <= MAX_TOKENS: raise ValueError("JADE's declared context limit is at most 8192 tokens") root = Path(model) if not root.is_dir(): if not revision or len(revision) != 40 or any(c not in "0123456789abcdefABCDEF" for c in revision): raise ValueError("Hub model revision must be an immutable 40-character commit") from huggingface_hub import snapshot_download root = Path(snapshot_download(repo_id=model, revision=revision)) manifest = validate_release(root) config = json.loads((root / "decision_config.json").read_text()) from transformers import AutoProcessor from vllm import LLM, SamplingParams from vllm.lora.request import LoRARequest self.processor = AutoProcessor.from_pretrained(config["base_model"], revision=config["revision"]) self.codes = config["codes"] self.token_ids = config["token_ids"] if ( len(self.codes) != MAX_OPTIONS or len(self.token_ids) != MAX_OPTIONS or len(set(self.token_ids)) != MAX_OPTIONS ): raise ValueError("Invalid decision vocabulary") actual = [self.processor.tokenizer.encode(code, add_special_tokens=False) for code in self.codes] if actual != [[token] for token in self.token_ids]: raise ValueError("Tokenizer differs from the trained decision vocabulary") self.temperature = float(config["temperature"]) if not math.isfinite(self.temperature) or self.temperature <= 0: raise ValueError("Checkpoint temperature must be positive and finite") self.limit = int(max_tokens) self.model_id = model self.SamplingParams = SamplingParams self.model = LLM( model=config["base_model"], revision=config["revision"], tokenizer_revision=config["revision"], dtype="bfloat16", max_model_len=self.limit, gpu_memory_utilization=float(gpu_memory_utilization), language_model_only=True, enable_lora=True, max_lora_rank=256, # The exported adapter includes the decision head. max_logprobs=MAX_OPTIONS, max_num_seqs=128, enable_prefix_caching=True, generation_config="vllm", ) self.adapter = LoRARequest("jade", 1, str(root / "vllm-adapter")) self.provenance = { "model": model, "revision": revision, "release_id": manifest["release_id"], "base_model": config["base_model"], "base_revision": config["revision"], "temperature": self.temperature, "context_limit_tokens": self.limit, "max_options": MAX_OPTIONS, "supported_types": ["choice", "noul"], "input_mode": "text", "truncation": False, "option_filtering": False, } def __call__(self, state: Any, questions: dict): if not questions: raise ValueError("Supply at least one question") prompts, settings, requests, keys_by_question = self._prepare(state, questions) results = self.model.generate(prompts, settings, lora_request=self.adapter, use_tqdm=False) gathered = {key: {} for key in questions} for (key, part), result in zip(requests, results, strict=True): returned = result.outputs[0].logprobs[0] for token in part: if token not in returned: raise ValueError(f"missing answer-token probability: {token}") gathered[key][token] = float(returned[token].logprob) answers = {} for key, question in questions.items(): keys = keys_by_question[key] logits = [gathered[key][token] for token in self.token_ids[:len(keys)]] probabilities = distribution(logits, self.temperature) if question["type"] == "noul": answers[key] = {"type": "noul", "noul": probabilities[1]} else: best = max(range(len(keys)), key=probabilities.__getitem__) answers[key] = { "type": "choice", "choice": keys[best], "probabilities": dict(zip(keys, probabilities, strict=True)), } return {"model": self.model_id, "answers": answers}, None def _prepare(self, state: Any, questions: dict): prompts = [] settings = [] requests = [] keys_by_question = {} for key, question in questions.items(): if question.get("type") not in {"choice", "noul"}: raise UnsupportedInput(f"unsupported question type: {question.get('type')}") keys, _ = options(question) if not 1 <= len(keys) <= MAX_OPTIONS: raise UnsupportedInput("question needs between 1 and 255 options") keys_by_question[key] = keys rendered = self.processor.apply_chat_template( decision_messages({"state": state, "question": question}, self.codes), tokenize=False, add_generation_prompt=True, enable_thinking=False, ) ids = self.processor.tokenizer(rendered, add_special_tokens=False)["input_ids"] if len(ids) + 1 > self.limit: raise UnsupportedInput(f"input needs {len(ids) + 1} tokens; capacity is {self.limit}; no truncation") # vLLM accepts at most 128 requested log-probability token IDs. # Each chunk scores the same full option set, so their values share # a normalizer and can be combined before temperature scaling. for offset in range(0, len(keys), LOGPROB_CHUNK_SIZE): part = self.token_ids[offset:min(offset + LOGPROB_CHUNK_SIZE, len(keys))] prompts.append({"prompt_token_ids": ids}) settings.append(self.SamplingParams( temperature=0, max_tokens=1, logprob_token_ids=part, allowed_token_ids=self.token_ids[:len(keys)], detokenize=False, )) requests.append((key, part)) return prompts, settings, requests, keys_by_question