JADE / jade /engine.py
theunnecessarythings's picture
Initial JADE release
fdbed23 verified
Raw History Blame Contribute Delete
8.24 kB
"""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