KillerKing93/Transformers-TextEngine-OpenAPI
0
1#!/usr/bin/env python2# -*- coding: utf-8 -*-3"""4Image Captioning Module for AI Marketplace Platform5 6Uses Salesforce/blip2-opt-2.7b for high-quality image descriptions.7Provides fast image captioning for multimodal chat functionality.8 9Usage:10 from image_caption import ImageCaptioner11 12 captioner = ImageCaptioner()13 caption = captioner.caption_image("path/to/image.jpg")14 print(f"Image description: {caption}")15"""16 17import os18import io19import torch20from PIL import Image21from transformers import AutoProcessor, Blip2ForConditionalGeneration22from typing import Optional, Union23import logging24 25# Setup logging26logging.basicConfig(level=logging.INFO)27logger = logging.getLogger(__name__)28 29class ImageCaptioner:30 """31 BLIP-2 Image Captioning with Salesforce/blip2-opt-2.7b32 33 Generates high-quality descriptions of images for multimodal AI interactions.34 """35 36 def __init__(self, model_name: str = "Salesforce/blip2-opt-2.7b"):37 self.model_name = model_name38 self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")39 self.model = None40 self.processor = None41 self._model_loaded = False42 43 def _load_model(self):44 """Load BLIP-2 model and processor on first use"""45 if self._model_loaded:46 return47 48 try:49 logger.info(f"Loading image captioning model: {self.model_name}")50 51 # Load processor and model52 self.processor = AutoProcessor.from_pretrained(self.model_name)53 54 # Load model with appropriate dtype55 if self.device.type == "cuda":56 self.model = Blip2ForConditionalGeneration.from_pretrained(57 self.model_name,58 torch_dtype=torch.float16,59 device_map="auto"60 )61 else:62 self.model = Blip2ForConditionalGeneration.from_pretrained(63 self.model_name,64 torch_dtype=torch.float3265 )66 self.model = self.model.to(self.device)67 68 # Set to eval mode69 self.model.eval()70 71 logger.info(f"Image captioning model loaded on {self.device}")72 self._model_loaded = True73 74 except Exception as e:75 logger.error(f"Failed to load image captioning model: {e}")76 raise77 78 def _prepare_image(self, image_input: Union[str, bytes, Image.Image]) -> Image.Image:79 """80 Prepare image for captioning81 82 Args:83 image_input: Path to image file, bytes, or PIL Image84 85 Returns:86 PIL Image87 """88 if isinstance(image_input, str):89 # Load from file path90 image = Image.open(image_input)91 elif isinstance(image_input, bytes):92 # Load from bytes93 image = Image.open(io.BytesIO(image_input))94 elif isinstance(image_input, Image.Image):95 # Already a PIL Image96 image = image_input97 else:98 raise ValueError(f"Unsupported image input type: {type(image_input)}")99 100 # Convert to RGB if needed101 if image.mode != "RGB":102 image = image.convert("RGB")103 104 return image105 106 def caption_image(107 self,108 image_input: Union[str, bytes, Image.Image],109 max_length: int = 50,110 num_beams: int = 5111 ) -> str:112 """113 Generate caption for image114 115 Args:116 image_input: Path to image file, bytes, or PIL Image117 max_length: Maximum caption length118 num_beams: Number of beams for beam search119 120 Returns:121 Generated caption as string122 """123 # Load model on first use124 self._load_model()125 126 try:127 # Prepare image128 image = self._prepare_image(image_input)129 130 # Process image131 inputs = self.processor(image, return_tensors="pt").to(self.device)132 133 # Generate caption134 with torch.no_grad():135 generated_ids = self.model.generate(136 **inputs,137 max_length=max_length,138 num_beams=num_beams,139 early_stopping=True,140 do_sample=False141 )142 143 # Decode caption144 caption = self.processor.batch_decode(145 generated_ids,146 skip_special_tokens=True147 )[0].strip()148 149 logger.info(f"Generated caption: {caption}")150 return caption151 152 except Exception as e:153 logger.error(f"Failed to generate caption: {e}")154 return "Unable to generate image description"155 156 def caption_with_context(157 self,158 image_input: Union[str, bytes, Image.Image],159 context: Optional[str] = None,160 max_length: int = 50,161 num_beams: int = 5162 ) -> str:163 """164 Generate caption with optional context/prompt165 166 Args:167 image_input: Path to image file, bytes, or PIL Image168 context: Optional context or prompt169 max_length: Maximum caption length170 num_beams: Number of beams for beam search171 172 Returns:173 Generated caption with context174 """175 # Load model on first use176 self._load_model()177 178 try:179 # Prepare image180 image = self._prepare_image(image_input)181 182 # Prepare prompt (optional)183 prompt = context if context else None184 185 # Process with prompt if provided186 if prompt:187 # For conditional captioning188 inputs = self.processor(189 image,190 text=prompt,191 return_tensors="pt"192 ).to(self.device)193 else:194 # For unconditional captioning195 inputs = self.processor(image, return_tensors="pt").to(self.device)196 197 # Generate caption198 with torch.no_grad():199 generated_ids = self.model.generate(200 **inputs,201 max_length=max_length,202 num_beams=num_beams,203 early_stopping=True,204 do_sample=False205 )206 207 # Decode caption208 caption = self.processor.batch_decode(209 generated_ids,210 skip_special_tokens=True211 )[0].strip()212 213 logger.info(f"Generated caption with context '{context}': {caption}")214 return caption215 216 except Exception as e:217 logger.error(f"Failed to generate caption with context: {e}")218 return "Unable to generate image description"219 220 def get_model_info(self) -> dict:221 """Get information about the loaded model"""222 return {223 "model_name": self.model_name,224 "device": str(self.device),225 "loaded": self._model_loaded,226 "model_type": "BLIP-2 Vision-Language Model"227 }228 229# Global instance for reuse230_image_captioner = None231 232def get_image_captioner() -> ImageCaptioner:233 """Get or create global image captioner instance"""234 global _image_captioner235 if _image_captioner is None:236 _image_captioner = ImageCaptioner()237 return _image_captioner238 239def caption_image(image_input: Union[str, bytes, Image.Image]) -> str:240 """241 Convenience function for image captioning242 243 Args:244 image_input: Path to image file, bytes, or PIL Image245 246 Returns:247 Generated caption as string248 """249 captioner = get_image_captioner()250 return captioner.caption_image(image_input)251 252def caption_image_with_context(253 image_input: Union[str, bytes, Image.Image],254 context: Optional[str] = None255) -> str:256 """257 Convenience function for image captioning with context258 259 Args:260 image_input: Path to image file, bytes, or PIL Image261 context: Optional context or prompt262 263 Returns:264 Generated caption as string265 """266 captioner = get_image_captioner()267 return captioner.caption_with_context(image_input, context)