""" Model management for Frame 0 Laboratory for MIA BAGEL 7B integration for advanced image analysis """ import logging import os import subprocess import spaces import torch from typing import Optional, Dict, Any, Tuple from PIL import Image from huggingface_hub import snapshot_download from accelerate import infer_auto_device_map, load_checkpoint_and_dispatch, init_empty_weights from config import ( BAGEL_CONFIG, get_device_config, get_bagel_device_map, BAGEL_PROMPTS, FLASH_ATTN_INSTALL ) from utils import clean_memory, safe_execute logger = logging.getLogger(__name__) class BaseImageAnalyzer: """Base class for image analysis models""" def __init__(self): self.model = None self.is_initialized = False self.device_config = get_device_config() def initialize(self) -> bool: """Initialize the model""" raise NotImplementedError def analyze_image(self, image: Image.Image) -> Tuple[str, Dict[str, Any]]: """Analyze image and return description""" raise NotImplementedError def cleanup(self) -> None: """Clean up model resources""" if hasattr(self, 'model') and self.model is not None: del self.model self.model = None clean_memory() class BagelAnalyzer(BaseImageAnalyzer): """BAGEL 7B model for advanced image analysis""" def __init__(self): super().__init__() self.inferencer = None self.tokenizer = None self.vae_model = None self.vae_transform = None self.vit_transform = None self._install_flash_attn() def _install_flash_attn(self): """Install flash attention dynamically""" try: logger.info("Installing flash attention...") result = subprocess.run( FLASH_ATTN_INSTALL["command"], env=FLASH_ATTN_INSTALL["env"], shell=FLASH_ATTN_INSTALL["shell"], capture_output=True, text=True ) if result.returncode == 0: logger.info("Flash attention installed successfully") else: logger.warning(f"Flash attention installation warning: {result.stderr}") except Exception as e: logger.warning(f"Flash attention installation failed: {e}") def _download_model(self) -> bool: """Download BAGEL model if not present""" try: logger.info("Downloading BAGEL model...") snapshot_download( cache_dir=BAGEL_CONFIG["cache_dir"], local_dir=BAGEL_CONFIG["local_model_path"], repo_id=BAGEL_CONFIG["model_repo"], local_dir_use_symlinks=False, resume_download=True, allow_patterns=BAGEL_CONFIG["download_patterns"], ) logger.info("BAGEL model downloaded successfully") return True except Exception as e: logger.error(f"BAGEL model download failed: {e}") return False def initialize(self) -> bool: """Initialize BAGEL model""" if self.is_initialized: return True try: # Download model if needed if not os.path.exists(BAGEL_CONFIG["local_model_path"]): if not self._download_model(): return False logger.info("Initializing BAGEL model...") # Import BAGEL components after flash attention installation from data.data_utils import add_special_tokens, pil_img2rgb from data.transforms import ImageTransform from inferencer import InterleaveInferencer from modeling.autoencoder import load_ae from modeling.bagel.qwen2_navit import NaiveCache from modeling.bagel import ( BagelConfig, Bagel, Qwen2Config, Qwen2ForCausalLM, SiglipVisionConfig, SiglipVisionModel ) from modeling.qwen2 import Qwen2Tokenizer model_path = BAGEL_CONFIG["local_model_path"] # Load configurations llm_config = Qwen2Config.from_json_file(os.path.join(model_path, "llm_config.json")) llm_config.qk_norm = True llm_config.tie_word_embeddings = False llm_config.layer_module = "Qwen2MoTDecoderLayer" vit_config = SiglipVisionConfig.from_json_file(os.path.join(model_path, "vit_config.json")) vit_config.rope = False vit_config.num_hidden_layers -= 1 # Load VAE self.vae_model, vae_config = load_ae(local_path=os.path.join(model_path, "ae.safetensors")) # Create BAGEL config config = BagelConfig( visual_gen=True, visual_und=True, llm_config=llm_config, vit_config=vit_config, vae_config=vae_config, vit_max_num_patch_per_side=70, connector_act='gelu_pytorch_tanh', latent_patch_size=2, max_latent_size=64, ) # Initialize model with empty weights with init_empty_weights(): language_model = Qwen2ForCausalLM(llm_config) vit_model = SiglipVisionModel(vit_config) self.model = Bagel(language_model, vit_model, config) self.model.vit_model.vision_model.embeddings.convert_conv2d_to_linear(vit_config, meta=True) # Load tokenizer self.tokenizer = Qwen2Tokenizer.from_pretrained(model_path) self.tokenizer, new_token_ids, _ = add_special_tokens(self.tokenizer) # Setup transforms vae_size = BAGEL_CONFIG["vae_transform_size"] vit_size = BAGEL_CONFIG["vit_transform_size"] self.vae_transform = ImageTransform(vae_size[0], vae_size[1], vae_size[2]) self.vit_transform = ImageTransform(vit_size[0], vit_size[1], vit_size[2]) # Setup device mapping device_map = infer_auto_device_map( self.model, max_memory={i: BAGEL_CONFIG["max_memory_per_gpu"] for i in range(torch.cuda.device_count())}, no_split_module_classes=["Bagel", "Qwen2MoTDecoderLayer"], ) # Apply custom device mapping for critical modules custom_mapping = get_bagel_device_map(self.device_config["gpu_count"]) device_map.update(custom_mapping) # Load model with checkpoints self.model = load_checkpoint_and_dispatch( self.model, checkpoint=os.path.join(model_path, "ema.safetensors"), device_map=device_map, offload_buffers=BAGEL_CONFIG["offload_buffers"], dtype=BAGEL_CONFIG["dtype"], force_hooks=BAGEL_CONFIG["force_hooks"], ).eval() # Initialize inferencer self.inferencer = InterleaveInferencer( model=self.model, vae_model=self.vae_model, tokenizer=self.tokenizer, vae_transform=self.vae_transform, vit_transform=self.vit_transform, new_token_ids=new_token_ids, ) self.is_initialized = True logger.info("BAGEL model initialized successfully") return True except Exception as e: logger.error(f"BAGEL initialization failed: {e}") self.cleanup() return False @spaces.GPU(duration=120) def analyze_image(self, image: Image.Image, prompt_type: str = "detailed_description") -> Tuple[str, Dict[str, Any]]: """Analyze image using BAGEL model""" if not self.is_initialized: success = self.initialize() if not success: return "BAGEL model not available", {"error": "Initialization failed"} try: # Get appropriate prompt system_prompt = BAGEL_PROMPTS.get(prompt_type, BAGEL_PROMPTS["detailed_description"]) # Prepare image for BAGEL if image.mode != 'RGB': image = image.convert('RGB') # Run inference through BAGEL logger.info("Running BAGEL inference...") # Use inferencer to analyze the image response = self.inferencer.inference_image_understanding( image=image, prompt=system_prompt, max_new_tokens=BAGEL_CONFIG["max_new_tokens"], temperature=BAGEL_CONFIG["temperature"], top_p=BAGEL_CONFIG["top_p"], do_sample=BAGEL_CONFIG["do_sample"] ) # Prepare metadata metadata = { "model": "BAGEL-7B", "device": self.device_config["device"], "confidence": 0.9, # BAGEL is highly reliable "prompt_type": prompt_type, "gpu_count": self.device_config.get("gpu_count", 1), "processing_mode": "GPU" if self.device_config["use_gpu"] else "CPU" } logger.info(f"BAGEL analysis complete: {len(response)} characters") return response, metadata except Exception as e: logger.error(f"BAGEL analysis failed: {e}") return "Analysis failed", {"error": str(e), "model": "BAGEL-7B"} def cleanup(self) -> None: """Clean up BAGEL resources""" try: if hasattr(self, 'inferencer') and self.inferencer is not None: del self.inferencer self.inferencer = None if hasattr(self, 'vae_model') and self.vae_model is not None: del self.vae_model self.vae_model = None super().cleanup() logger.info("BAGEL resources cleaned up") except Exception as e: logger.warning(f"BAGEL cleanup warning: {e}") class FallbackAnalyzer(BaseImageAnalyzer): """Simple fallback analyzer when BAGEL is not available""" def __init__(self): super().__init__() def initialize(self) -> bool: """Fallback is always ready""" self.is_initialized = True return True def analyze_image(self, image: Image.Image) -> Tuple[str, Dict[str, Any]]: """Provide basic image description""" try: # Basic image analysis width, height = image.size mode = image.mode # Simple descriptive text based on image properties aspect_ratio = width / height if aspect_ratio > 1.5: orientation = "landscape" elif aspect_ratio < 0.75: orientation = "portrait" else: orientation = "square" description = f"A {orientation} photograph with {mode} color mode, {width}x{height} pixels. Professional image suitable for detailed analysis and prompt generation." metadata = { "model": "Fallback", "device": "cpu", "confidence": 0.5, "image_size": f"{width}x{height}", "color_mode": mode, "orientation": orientation } return description, metadata except Exception as e: logger.error(f"Fallback analysis failed: {e}") return "Basic image detected", {"error": str(e), "model": "Fallback"} class ModelManager: """Manager for handling image analysis models""" def __init__(self, preferred_model: str = "bagel"): self.preferred_model = preferred_model self.analyzers = {} self.current_analyzer = None def get_analyzer(self, model_name: str = None) -> Optional[BaseImageAnalyzer]: """Get or create analyzer for specified model""" model_name = model_name or self.preferred_model if model_name not in self.analyzers: if model_name == "bagel": self.analyzers[model_name] = BagelAnalyzer() elif model_name == "fallback": self.analyzers[model_name] = FallbackAnalyzer() else: logger.warning(f"Unknown model: {model_name}, using fallback") model_name = "fallback" self.analyzers[model_name] = FallbackAnalyzer() return self.analyzers[model_name] def analyze_image(self, image: Image.Image, model_name: str = None) -> Tuple[str, Dict[str, Any]]: """Analyze image with specified or preferred model""" # Try preferred model first analyzer = self.get_analyzer(model_name) if analyzer is None: return "No analyzer available", {"error": "Model not found"} success, result = safe_execute(analyzer.analyze_image, image) if success and result[1].get("error") is None: return result else: # Fallback to simple analyzer if main model fails logger.warning(f"Primary model failed, using fallback: {result}") fallback_analyzer = self.get_analyzer("fallback") fallback_success, fallback_result = safe_execute(fallback_analyzer.analyze_image, image) if fallback_success: return fallback_result else: return "All analyzers failed", {"error": "Complete analysis failure"} def cleanup_all(self) -> None: """Clean up all model resources""" for analyzer in self.analyzers.values(): analyzer.cleanup() self.analyzers.clear() clean_memory() logger.info("All analyzers cleaned up") # Global model manager instance model_manager = ModelManager(preferred_model="bagel") def analyze_image(image: Image.Image, model_name: str = None) -> Tuple[str, Dict[str, Any]]: """ Convenience function for image analysis using BAGEL Args: image: PIL Image to analyze model_name: Optional model name ("bagel" or "fallback") Returns: Tuple of (description, metadata) """ return model_manager.analyze_image(image, model_name) # Export main components __all__ = [ "BaseImageAnalyzer", "BagelAnalyzer", "FallbackAnalyzer", "ModelManager", "model_manager", "analyze_image" ]