Team Ai
Apppublic

KillerKing93/Transformers-TextEngine-OpenAPI

sourceHugging Faceupdated 11mo agoView on Hugging Face
0likes
image_caption.py267 linesDownload Raw Back to root
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)