SPX-CD-Pro / infer.py
Surd-AI's picture
Initial model release
04f53c5
Raw History Blame Contribute Delete
7.62 kB
"""Inference CLI. Reference answers are never passed to the model."""
import argparse
import json
import math
import os
from pathlib import Path
from decision import decide, validate_labels
RELEASE = json.loads(Path(__file__).with_name("release.json").read_text())
MAX_PIXELS = 512 * 32 * 32
class TransformersBackend:
def __init__(self, base, adapter, max_context):
import torch
from transformers import AutoProcessor, AutoModelForMultimodalLM
from peft import PeftModel
self.torch = torch
self.processor = AutoProcessor.from_pretrained(base, max_pixels=MAX_PIXELS)
self.model = AutoModelForMultimodalLM.from_pretrained(
base, dtype=torch.bfloat16, device_map="auto")
self.model = PeftModel.from_pretrained(self.model, adapter).eval()
self.max_context = max_context
def logits(self, prompt, images, allowed):
if images:
inputs = self.processor(text=[prompt], images=list(images), return_tensors="pt")
else:
inputs = self.processor.tokenizer(prompt, add_special_tokens=False, return_tensors="pt")
if inputs["input_ids"].shape[1] + 1 > self.max_context:
raise ValueError("Prompt exceeds context limit; no automatic truncation")
inputs = inputs.to(self.model.get_input_embeddings().weight.device)
with self.torch.inference_mode():
outputs = self.model(**inputs, logits_to_keep=1, use_cache=False)
# Preserve the model's native final logits.
return outputs.logits[0, -1, allowed].float().cpu().tolist()
class VLLMBackend:
def __init__(self, base, adapter, max_context, gpu_memory):
from transformers import AutoProcessor
from vllm import LLM, SamplingParams
import vllm.sampling_params as sampling_module
from vllm.lora.request import LoRARequest
# This runner requests only the candidate vocabulary, up to 1,024 IDs.
sampling_module.MAX_LOGPROB_TOKEN_IDS = 1025
cfg = json.loads(Path(adapter, "adapter_config.json").read_text())
self.processor = AutoProcessor.from_pretrained(base, max_pixels=MAX_PIXELS)
self.params = SamplingParams
self.max_context = max_context
self.model = LLM(
model=base, dtype="bfloat16", enable_lora=True, max_lora_rank=int(cfg["r"]),
max_model_len=max_context, max_num_seqs=2,
max_num_batched_tokens=2048, gpu_memory_utilization=gpu_memory,
enforce_eager=True, logprobs_mode="raw_logits", max_logprobs=1025,
attention_config={"backend": "TRITON_ATTN"},
limit_mm_per_prompt={"image": 5, "video": 0, "audio": 0},
mm_processor_kwargs={"max_pixels": MAX_PIXELS}, seed=20261004)
self.request = LoRARequest(RELEASE["model_id"], 1, str(Path(adapter).resolve()))
def logits(self, prompt, images, allowed):
length = (self.processor(text=[prompt], images=list(images), return_tensors="pt")
["input_ids"].shape[1] if images else
len(self.processor.tokenizer.encode(prompt, add_special_tokens=False)))
if length + 1 > self.max_context:
raise ValueError("Prompt exceeds context limit; no automatic truncation")
inp = {"prompt": prompt}
if images:
inp["multi_modal_data"] = {"image": list(images)}
out = self.model.generate(
[inp], self.params(temperature=0, max_tokens=1,
allowed_token_ids=allowed, logprob_token_ids=allowed),
lora_request=self.request, use_tqdm=False)[0]
scores = out.outputs[0].logprobs[0]
return [scores[token_id].logprob for token_id in allowed]
def main():
p = argparse.ArgumentParser(description=__doc__)
p.add_argument("--backend", choices=["transformers", "vllm"], default="transformers")
p.add_argument("--base", default=RELEASE["base_model"])
p.add_argument("--adapter", required=True)
p.add_argument("--input", required=True, help="One JSON object or a JSONL file")
p.add_argument("--output", required=True, help="New JSONL file; existing files are not overwritten")
p.add_argument("--image-root", default=None)
p.add_argument("--labels", default=str(Path(__file__).with_name("label_tokens.json")))
p.add_argument("--effort", type=int, choices=range(1, 6), default=1,
help="Budget of distinct candidate orderings (1–5)")
p.add_argument("--prompt-format", choices=["cmdb", "open-format"], default="cmdb")
p.add_argument("--temperature", type=float, default=1.0)
p.add_argument("--max-context", type=int, default=65536, help="Total context token budget (default: 65536)")
p.add_argument("--gpu-memory", type=float, default=0.8)
p.add_argument("--cache-env", help="Optional JSON environment map, applied before ML imports")
p.add_argument("--dry-run", action="store_true", help="Validate inputs without loading a model")
a = p.parse_args()
if not math.isfinite(a.temperature) or a.temperature <= 0:
p.error("--temperature must be positive and finite")
if a.max_context < 2 or not 0 < a.gpu_memory < 1:
p.error("Invalid context limit or GPU memory fraction")
if a.cache_env:
os.environ.update({k: str(v) for k, v in json.loads(Path(a.cache_env).read_text()).items()})
inp = Path(a.input)
if inp.suffix == ".json":
cases = [json.loads(inp.read_text())]
else:
cases = [json.loads(line) for line in inp.open() if line.strip()]
from decision import normalize_case
labels = json.loads(Path(a.labels).read_text())
image_root = Path(a.image_root) if a.image_root else inp.parent
ids = set()
for case in cases:
case["_prompt_format"] = a.prompt_format
normalized = normalize_case(case)
if a.prompt_format == "open-format" and (case.get("image_refs") or normalized["kind"] == "multi_choice"):
raise ValueError("Use --prompt-format cmdb for images or multi-select")
if normalized["id"] in ids:
raise ValueError("Input IDs must be unique")
ids.add(normalized["id"])
if len(normalized["options"]) > len(labels["ids"]):
raise ValueError("Too many candidates")
for ref in case.get("image_refs", []):
path = (image_root / ref).resolve()
if not path.is_relative_to(image_root.resolve()) or not path.is_file():
raise ValueError(f"Image must exist within image-root: {ref}")
if a.dry_run:
print(json.dumps({"validated_inputs": len(cases), "label_count": len(labels["ids"])}))
return
if Path(a.output).exists():
raise FileExistsError("Choose a new output file")
backend = (TransformersBackend(a.base, a.adapter, a.max_context) if a.backend == "transformers"
else VLLMBackend(a.base, a.adapter, a.max_context, a.gpu_memory))
validate_labels(backend.processor, labels)
from PIL import Image
with Path(a.output).open("x") as out:
for case in cases:
images = []
for ref in case.get("image_refs", []):
with Image.open(image_root / ref) as image:
images.append(image.convert("RGB"))
if len(images) > 5:
raise ValueError("At most five images are supported by this runner")
result = decide(case, backend, labels, images, a.effort, a.temperature)
out.write(json.dumps(result, ensure_ascii=False) + "\n")
out.flush()
if __name__ == "__main__":
main()