TestSpace3 / app.py
ProfRom's picture
Saar - Sanity Check
523547d verified
Raw
History Blame Contribute Delete
3.17 kB
import torch
import gradio as gr
from PIL import Image
from transformers import (
BlipProcessor,
BlipForConditionalGeneration,
pipeline
)
# Select device
device = "cuda" if torch.cuda.is_available() else "cpu"
torch_dtype = torch.float16 if torch.cuda.is_available() else torch.float32
# Load BLIP captioning model directly
processor = BlipProcessor.from_pretrained("Salesforce/blip-image-captioning-base")
caption_model = BlipForConditionalGeneration.from_pretrained(
"Salesforce/blip-image-captioning-base",
torch_dtype=torch_dtype
).to(device)
# Load image classification model
classifier = pipeline(
task="image-classification",
model="google/vit-base-patch16-224",
device=0 if torch.cuda.is_available() else -1
)
print("Models loaded successfully.")
def generate_caption(image):
inputs = processor(images=image, return_tensors="pt").to(device, torch_dtype)
output = caption_model.generate(**inputs, max_new_tokens=30)
caption = processor.decode(output[0], skip_special_tokens=True)
return caption
def build_summary(caption: str, label: str) -> str:
caption = caption.strip() if caption else "No caption available"
label = label.strip() if label else "unknown object"
return (
f"The captioning model describes the image as: {caption}. "
f"The image classification model identifies the main subject as: {label}. "
f"Taken together, the image appears to focus on this subject or scene."
)
def analyze_image(image):
try:
if image is None:
return (
"Please upload an image.",
"No classification available.",
"Please upload an image first."
)
image = image.convert("RGB")
# Captioning
caption = generate_caption(image)
print("CAPTION RESULT:", caption)
# Classification
class_result = classifier(image)
print("CLASSIFICATION RESULT:", class_result)
if isinstance(class_result, list) and len(class_result) > 0:
top_label = class_result[0].get("label", "Unknown")
top_score = class_result[0].get("score", 0.0)
classification_text = f"{top_label} (confidence: {top_score:.4f})"
else:
top_label = "Unknown"
classification_text = "No classification generated."
summary = build_summary(caption, top_label)
return caption, classification_text, summary
except Exception as e:
print("ERROR:", str(e))
error_text = f"Error: {str(e)}"
return error_text, error_text, error_text
demo = gr.Interface(
fn=analyze_image,
inputs=gr.Image(type="pil", label="Upload an Image"),
outputs=[
gr.Textbox(label="Generated Caption"),
gr.Textbox(label="Top Classification"),
gr.Textbox(label="Combined Summary", lines=4)
],
title="Image Captioning, Classification, and Summary App",
description=(
"Upload an image to generate an automatic caption, predict the main image class, "
"and produce a short combined summary."
),
)
if __name__ == "__main__":
demo.launch()