Instructions to use theunnecessarythings/JADE with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use theunnecessarythings/JADE with PEFT:
Task type is invalid.
- Notebooks
- Google Colab
- Kaggle
Download jade/engine.py from theunnecessarythings/JADE: direct link, hf CLI and curl.
- Browser
- Download file 8.24 kB
-
https://huggingface.co/theunnecessarythings/JADE/resolve/main/jade/engine.py
- Command line
-
hf download hf://theunnecessarythings/JADE/jade/engine.py
-
curl -L -o engine.py https://huggingface.co/theunnecessarythings/JADE/resolve/main/jade/engine.py
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 | |