clef-NVFP4 / clef_vllm.py
simonlehmann's picture
Add files using upload-large-folder tool
817ac58 verified
Raw History Blame Contribute Delete
6.23 kB
"""Run Clef-NVFP4 on vLLM: the backbone as a pooling model (final hidden state of every
token, on NVFP4/FP8 kernels), Clef's joint schema head on top in the same process.
from clef_vllm import ClefVLLM
clef = ClefVLLM("path/to/clef-NVFP4")
clef.systemone({"model": "clef", "state": "...", "questions": {...}})
Tested with vLLM 0.23.1 (nightly, CUDA 13, aarch64 / Blackwell sm_121).
"""
from __future__ import annotations
import json
import os
import sys
from pathlib import Path
from typing import Any
# The engine runs in this process: no __main__ guard needed, no IPC per decision.
os.environ.setdefault("VLLM_ENABLE_V1_MULTIPROCESSING", "0")
import torch # noqa: E402
from safetensors import safe_open # noqa: E402
from safetensors.torch import load_file # noqa: E402
sys.path.insert(0, str(Path(__file__).resolve().parent))
from joint_schema_model import ( # noqa: E402
QUESTION_TYPES,
JointSchemaHead,
collate_records,
encode_record,
systemone_answer,
)
class ClefVLLM:
def __init__(self, path: str | Path, max_model_len: int = 16384, gpu_memory_utilization: float = 0.4,
max_images: int = 1, max_videos: int = 0, **llm_kwargs: Any) -> None:
from transformers import AutoProcessor
from vllm import LLM
from vllm.config import PoolerConfig
path = Path(path)
self.processor = AutoProcessor.from_pretrained(path)
self.tokenizer = self.processor.tokenizer
self.media_pads = {self.tokenizer.convert_tokens_to_ids(t) for t in ("<|image_pad|>", "<|video_pad|>")}
self.llm = LLM(
model=str(path),
runner="pooling",
pooler_config=PoolerConfig(task="token_embed", tok_pooling_type="ALL", use_activation=False),
max_model_len=max_model_len,
gpu_memory_utilization=gpu_memory_utilization,
limit_mm_per_prompt={"image": max_images, "video": max_videos},
enable_prefix_caching=False,
**llm_kwargs,
)
self.device = torch.device("cuda")
head = JointSchemaHead(**json.loads((path / "joint_head_config.json").read_text()))
head.load_state_dict(load_file(path / "joint_head.safetensors"), strict=True)
self.head = head.to(self.device, torch.bfloat16).eval()
# The head reads lm_head rows as option embeddings; vLLM's pooling model drops lm_head.
weight_map = json.loads((path / "model.safetensors.index.json").read_text())["weight_map"]
with safe_open(str(path / weight_map["lm_head.weight"]), framework="pt", device="cuda") as f:
self.lm_head = f.get_tensor("lm_head.weight").to(torch.bfloat16)
def _prompt(self, record: dict[str, Any], input_ids: tuple[int, ...]) -> dict[str, Any]:
# Collapse each expanded image/video pad run to one token; vLLM re-expands it from the
# media, so the returned hidden states line up with Clef's input_ids position for position.
ids: list[int] = []
for i, token in enumerate(input_ids):
if token in self.media_pads and i and input_ids[i - 1] == token:
continue
ids.append(token)
prompt: dict[str, Any] = {"prompt_token_ids": ids}
media = {}
if record.get("images"):
media["image"] = record["images"] if len(record["images"]) > 1 else record["images"][0]
if record.get("videos"):
media["video"] = record["videos"] if len(record["videos"]) > 1 else record["videos"][0]
if media:
prompt["multi_modal_data"] = media
return prompt
@torch.inference_mode()
def probabilities(self, records: list[dict[str, Any]], max_length: int = 16384) -> list[dict[str, dict[str, float]]]:
"""Per record, per question: {option_id: probability}. Records are batched through vLLM."""
encoded = [encode_record(self.tokenizer, r, max_length=max_length, processor=self.processor) for r in records]
outputs = self.llm.encode([self._prompt(r, e.input_ids) for r, e in zip(records, encoded)],
pooling_task="token_embed", use_tqdm=False)
results = []
for enc, out in zip(encoded, outputs):
hidden = out.outputs.data.to(self.device, torch.bfloat16)
if hidden.shape[0] != len(enc.input_ids):
raise RuntimeError(f"vLLM returned {hidden.shape[0]} positions for {len(enc.input_ids)} input tokens")
batch = collate_records([enc], self.tokenizer.pad_token_id, self.device)
logits = self.head(hidden.unsqueeze(0), batch["input_ids"], batch["attention_mask"],
batch["records"], self.lm_head)[0]
results.append({
q.question_id: dict(zip(q.option_ids, l.float().softmax(-1).tolist()))
for q, l in zip(enc.questions, logits)
})
return results
def systemone(self, request: dict[str, Any], max_length: int = 16384) -> dict[str, Any]:
"""Same request/response body as joint_schema_model.systemone (Jev/SystemOne /v1/systemone)."""
questions = request.get("questions")
if not isinstance(request.get("model"), str) or "state" not in request:
raise ValueError("model and state are required")
if not isinstance(questions, dict) or not questions:
raise ValueError("at least one question is required")
for question_id, question in questions.items():
if question.get("type") not in QUESTION_TYPES:
raise ValueError(f"{question_id}: type must be noul, choice, or score")
if question["type"] != "noul" and not question.get("criteria"):
raise ValueError(f"{question_id}: criteria must not be empty")
probs = self.probabilities([request], max_length=max_length)[0]
encoded_len = len(encode_record(self.tokenizer, request, max_length=max_length, processor=self.processor).input_ids)
return {
"model": request["model"],
"answers": {qid: systemone_answer(questions[qid], p) for qid, p in probs.items()},
"usage": {"input_tokens": encoded_len, "output_tokens": 0},
}