#!/usr/bin/env python3 """Single-image inference with strict Miril-DroneVLM-2B-2 response validation.""" from __future__ import annotations import argparse import json import subprocess import sys from pathlib import Path from typing import Any from PIL import Image from router_contract import ( TYPED_JSON_ROUTER_SYSTEM_PROMPT, drawable_point, typed_response_errors, ) def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--image", required=True) query = parser.add_mutually_exclusive_group(required=True) query.add_argument("--prompt") query.add_argument("--audio", help="Experimental spoken-question audio file.") parser.add_argument( "--audio-transcript", help="Optional audit/display transcript; it is not supplied to the model.", ) parser.add_argument("--audio-sample-rate", type=int, default=16000) parser.add_argument("--max-audio-seconds", type=float, default=30.0) parser.add_argument("--model-id", default="MirilAI/Miril-DroneVLM-2B-2") parser.add_argument("--processor-id") parser.add_argument("--max-new-tokens", type=int, default=384) parser.add_argument("--load-4bit", action="store_true") parser.add_argument("--output-json") return parser.parse_args() def load_audio_mono_float32( path: str | Path, *, sample_rate: int = 16000, max_seconds: float = 30.0, ) -> Any: import numpy as np audio_path = Path(path).expanduser() if not audio_path.is_file(): raise FileNotFoundError(f"audio file does not exist: {audio_path}") if sample_rate <= 0 or max_seconds <= 0: raise ValueError("audio sample rate and maximum duration must be positive") process = subprocess.run( [ "ffmpeg", "-hide_banner", "-loglevel", "error", "-i", str(audio_path), "-f", "f32le", "-acodec", "pcm_f32le", "-ac", "1", "-ar", str(sample_rate), "pipe:1", ], check=True, stdout=subprocess.PIPE, ) waveform = np.frombuffer(process.stdout, dtype=" max_samples: raise ValueError( f"audio exceeds --max-audio-seconds ({waveform.size / sample_rate:.2f}s)" ) return waveform def parse_bare_json(text: str) -> tuple[dict[str, Any] | None, list[str]]: stripped = text.strip() try: payload = json.loads(stripped) except Exception as exc: return None, [f"invalid bare JSON: {exc}"] if not isinstance(payload, dict): return None, ["response must be a JSON object"] errors = typed_response_errors(payload) return (payload if not errors else None), errors def dispatch_summary(payload: dict[str, Any]) -> dict[str, Any]: response_type = payload["type"] summary: dict[str, Any] = { "type": response_type, "dispatch_key": response_type, "text": payload.get("caption") if response_type != "answer" else payload["answer"], "drawable": False, "trackable": False, } if response_type == "location": summary["dispatch_key"] = f"location:{payload['intent']}" elif response_type == "pointing": summary["dispatch_key"] = f"pointing:{payload['action']}" point = drawable_point(payload) if point is not None: x, y, visual_mode = point summary.update( { "drawable": True, "visual_mode": visual_mode, "normalized_xy": [x, y], "trackable": visual_mode == "precise_target", } ) return summary def load_model(model_id: str, processor_id: str, load_4bit: bool) -> tuple[Any, Any]: import torch import transformers from transformers import AutoProcessor processor = AutoProcessor.from_pretrained(processor_id) model_class = None for class_name in ( "AutoModelForMultimodalLM", "AutoModelForImageTextToText", "AutoModelForVision2Seq", ): model_class = getattr(transformers, class_name, None) if model_class is not None: break if model_class is None: raise RuntimeError( "Installed transformers has no supported multimodal auto-model class" ) model_kwargs: dict[str, Any] = { "device_map": "auto", "dtype": torch.bfloat16 if torch.cuda.is_available() else torch.float32, } if load_4bit: from transformers import BitsAndBytesConfig model_kwargs["quantization_config"] = BitsAndBytesConfig( load_in_4bit=True, bnb_4bit_quant_type="nf4", bnb_4bit_compute_dtype=torch.bfloat16, ) model = model_class.from_pretrained(model_id, **model_kwargs) model.eval() return model, processor def generate( model: Any, processor: Any, image: Image.Image, prompt: str | None, *, audio_waveform: Any | None = None, max_new_tokens: int, ) -> str: import torch if (prompt is None) == (audio_waveform is None): raise ValueError("provide exactly one of prompt or audio_waveform") user_content: list[dict[str, Any]] = [{"type": "image", "image": image}] if audio_waveform is not None: user_content.append({"type": "audio", "audio": audio_waveform}) else: user_content.append({"type": "text", "text": prompt}) messages = [ { "role": "system", "content": [{"type": "text", "text": TYPED_JSON_ROUTER_SYSTEM_PROMPT}], }, { "role": "user", "content": user_content, }, ] rendered = processor.apply_chat_template( messages, tokenize=False, add_generation_prompt=True, ) processor_kwargs: dict[str, Any] = { "text": [rendered], "images": [[image]], "return_tensors": "pt", "padding": True, } if audio_waveform is not None: processor_kwargs["audio"] = [audio_waveform] inputs = processor(**processor_kwargs) device = next(model.parameters()).device inputs = inputs.to(device) with torch.no_grad(): output = model.generate( **inputs, max_new_tokens=max_new_tokens, do_sample=False, num_beams=1, ) generated = output[:, inputs["input_ids"].shape[-1] :] return processor.batch_decode(generated, skip_special_tokens=True)[0].strip() def main() -> int: args = parse_args() if args.audio_transcript and not args.audio: raise ValueError("--audio-transcript requires --audio") processor_id = args.processor_id or args.model_id model, processor = load_model(args.model_id, processor_id, args.load_4bit) image = Image.open(args.image).convert("RGB") audio_waveform = ( load_audio_mono_float32( args.audio, sample_rate=args.audio_sample_rate, max_seconds=args.max_audio_seconds, ) if args.audio else None ) raw = generate( model, processor, image, args.prompt, audio_waveform=audio_waveform, max_new_tokens=args.max_new_tokens, ) payload, errors = parse_bare_json(raw) result = { "ok": payload is not None, "model_id": args.model_id, "input_modality": "image_audio" if args.audio else "image_text", "prompt": args.prompt, "audio_path": str(Path(args.audio).expanduser()) if args.audio else None, "audio_transcript_for_audit_only": args.audio_transcript, "raw": raw, "payload": payload, "validation_errors": errors, "dispatch": dispatch_summary(payload) if payload is not None else None, } rendered = json.dumps(result, indent=2, ensure_ascii=False) print(rendered) if args.output_json: output_path = Path(args.output_json) output_path.parent.mkdir(parents=True, exist_ok=True) output_path.write_text(rendered + "\n", encoding="utf-8") if errors: print( "Response rejected; do not dispatch or draw coordinates.", file=sys.stderr ) return 2 return 0 if __name__ == "__main__": raise SystemExit(main())