from pathlib import Path from time import perf_counter import torch from transformers import ( AutoProcessor, Qwen3VLForConditionalGeneration, ) from qwen_vl_utils import process_vision_info from cockpit_grounding.models.base import InferenceTiming class Qwen3VLGrounder: def __init__( self, model_path: str, ) -> None: local_model_path = Path(model_path).expanduser().resolve() if not local_model_path.is_dir(): raise FileNotFoundError( f"Local model directory not found: {local_model_path}" ) if not torch.cuda.is_available(): raise RuntimeError("Qwen3VLGrounder requires a CUDA device") self.model_path = str(local_model_path) print("[Model] Loading local model:") print(self.model_path) self.model = ( Qwen3VLForConditionalGeneration .from_pretrained( self.model_path, dtype=torch.bfloat16, device_map={"": 0}, local_files_only=True, ) ) self.processor = AutoProcessor.from_pretrained( self.model_path, local_files_only=True, ) self.model.eval() print("[Model] Loaded successfully") print( "[Model] GPU:", torch.cuda.get_device_name(self.model.device) ) @torch.inference_mode() def generate( self, image_path: str, prompt: str, max_new_tokens: int = 128, ) -> str: raw_output, _, _ = self.generate_with_metrics( image_path=image_path, prompt=prompt, max_new_tokens=max_new_tokens, ) return raw_output @torch.inference_mode() def generate_with_metrics( self, image_path: str, prompt: str, max_new_tokens: int = 128, ) -> tuple[str, InferenceTiming, dict[str, float]]: """Generate a response and report synchronized stage timings.""" image_path = Path(image_path).resolve() if not image_path.exists(): raise FileNotFoundError(image_path) self._synchronize_cuda() torch.cuda.reset_peak_memory_stats(self.model.device) total_start = perf_counter() preprocess_start = total_start messages = [ { "role": "user", "content": [ { "type": "image", "image": image_path.as_uri(), }, { "type": "text", "text": prompt, }, ], } ] text = self.processor.apply_chat_template( messages, tokenize=False, add_generation_prompt=True, ) images, videos, video_kwargs = process_vision_info( messages, image_patch_size=16, return_video_kwargs=True, return_video_metadata=True, ) if videos is not None: videos, video_metadatas = zip(*videos) videos = list(videos) video_metadatas = list(video_metadatas) else: video_metadatas = None inputs = self.processor( text=text, images=images, videos=videos, video_metadata=video_metadatas, return_tensors="pt", do_resize=False, **video_kwargs, ) inputs = inputs.to(self.model.device) self._synchronize_cuda() preprocess_ms = (perf_counter() - preprocess_start) * 1000.0 generate_start = perf_counter() generated_ids = self.model.generate( **inputs, max_new_tokens=max_new_tokens, do_sample=False, ) self._synchronize_cuda() generate_ms = (perf_counter() - generate_start) * 1000.0 decode_start = perf_counter() generated_ids_trimmed = [ output_ids[len(input_ids):] for input_ids, output_ids in zip( inputs.input_ids, generated_ids, ) ] output_text = self.processor.batch_decode( generated_ids_trimmed, skip_special_tokens=True, clean_up_tokenization_spaces=False, )[0] self._synchronize_cuda() decode_ms = (perf_counter() - decode_start) * 1000.0 total_ms = (perf_counter() - total_start) * 1000.0 peak_cuda_memory_mb = self._bytes_to_mb( torch.cuda.max_memory_allocated(self.model.device) ) timing = InferenceTiming( preprocess_ms=preprocess_ms, generate_ms=generate_ms, decode_ms=decode_ms, total_ms=total_ms, ) extra_metrics = { "peak_cuda_memory_mb": peak_cuda_memory_mb, } return output_text, timing, extra_metrics @staticmethod def _bytes_to_mb(value: int) -> float: return value / (1024.0 * 1024.0) def cuda_memory_metrics(self) -> dict[str, float]: return { "model_cuda_allocated_mb": self._bytes_to_mb( torch.cuda.memory_allocated(self.model.device) ), "model_cuda_reserved_mb": self._bytes_to_mb( torch.cuda.memory_reserved(self.model.device) ), } def implementation_metadata(self) -> dict[str, str]: return { "transformers_model_class": type(self.model).__name__, "transformers_processor_class": type(self.processor).__name__, } def _synchronize_cuda(self) -> None: torch.cuda.synchronize(self.model.device)