Phramer_AI / models.py
Malaji71's picture
Update models.py
8d6efc2 verified
Raw History Blame
15 kB
"""
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"
]