from pathlib import Path from time import perf_counter import torch from transformers import AutoModelForMultimodalLM, AutoProcessor from cockpit_grounding.models.base import InferenceTiming class Qwen35Grounder: 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("Qwen35Grounder requires a CUDA device") self.model_path = str(local_model_path) print("[Model] Loading local Qwen3.5 model:") print(self.model_path) self.model = AutoModelForMultimodalLM.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(f"[Model] Class: {type(self.model).__name__}") 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_text( self, prompt: str, max_new_tokens: int = 256, ) -> str: messages = [ { "role": "user", "content": [{"type": "text", "text": prompt}], } ] inputs = self.processor.apply_chat_template( messages, add_generation_prompt=True, tokenize=True, return_dict=True, return_tensors="pt", ) inputs = inputs.to(self.model.device) output_ids = self.model.generate( **inputs, max_new_tokens=max_new_tokens, do_sample=False, ) generated_ids = output_ids[:, inputs["input_ids"].shape[-1] :] return self.processor.batch_decode( generated_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False, )[0] @torch.inference_mode() def generate_with_metrics( self, image_path: str, prompt: str, max_new_tokens: int = 128, ) -> tuple[str, InferenceTiming, dict[str, float]]: image = Path(image_path).expanduser().resolve() if not image.is_file(): raise FileNotFoundError(image) 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", "path": str(image)}, {"type": "text", "text": prompt}, ], } ] inputs = self.processor.apply_chat_template( messages, add_generation_prompt=True, tokenize=True, return_dict=True, return_tensors="pt", ) inputs = inputs.to(self.model.device) self._synchronize_cuda() preprocess_ms = (perf_counter() - preprocess_start) * 1000.0 generate_start = perf_counter() output_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 = output_ids[:, inputs["input_ids"].shape[-1] :] output_text = self.processor.batch_decode( generated_ids, 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 timing = InferenceTiming( preprocess_ms=preprocess_ms, generate_ms=generate_ms, decode_ms=decode_ms, total_ms=total_ms, ) metrics = { "peak_cuda_memory_mb": self._bytes_to_mb( torch.cuda.max_memory_allocated(self.model.device) ) } return output_text, timing, metrics 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__, } @staticmethod def _bytes_to_mb(value: int) -> float: return value / (1024.0 * 1024.0) def _synchronize_cuda(self) -> None: torch.cuda.synchronize(self.model.device)