"""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()