CoolFace
Apppublic

zanemotiwala/image-recognition-caption

sourceHugging Faceunknownupdated 1y agoView on Hugging Face
0likes
app.py65 linesDownload Raw Back to root
1import gradio as gr2from transformers import AutoTokenizer, AutoImageProcessor, VisionEncoderDecoderModel, ViTForImageClassification3import torch4 5# Set device6device = "cuda" if torch.cuda.is_available() else "cpu"7 8# Model checkpoints9caption_model_ckpt = "nlpconnect/vit-gpt2-image-captioning"10classify_model_ckpt = "google/vit-base-patch16-224"11 12# Load captioning components13tokenizer = AutoTokenizer.from_pretrained(caption_model_ckpt)14image_processor = AutoImageProcessor.from_pretrained(caption_model_ckpt)15caption_model = VisionEncoderDecoderModel.from_pretrained(caption_model_ckpt).to(device)16 17# Load classification model18classify_processor = AutoImageProcessor.from_pretrained(classify_model_ckpt)19classification_model = ViTForImageClassification.from_pretrained(classify_model_ckpt).to(device)20 21# Captioning function22def get_caption(image):23    if image is None:24        return "No image uploaded."25 26    image = image.convert("RGB")27    pixel_values = image_processor(images=image, return_tensors="pt").pixel_values.to(device)28    output_ids = caption_model.generate(pixel_values, max_length=64, num_beams=4)[0]29    caption = tokenizer.decode(output_ids, skip_special_tokens=True)30    return caption31 32# Classification function33def classify_image(image):34    if image is None:35        return {"Error": "No image uploaded."}36    37    image = image.convert("RGB")38    inputs = classify_processor(images=image, return_tensors="pt").to(device)39    outputs = classification_model(**inputs)40    probs = torch.nn.functional.softmax(outputs.logits, dim=-1)41    top_probs, top_labels = torch.topk(probs, 5)42 43    results = {44        classification_model.config.id2label[label.item()]: round(prob.item(), 4)45        for label, prob in zip(top_labels[0], top_probs[0])46    }47    return results48 49# Gradio app50with gr.Blocks(title="Image Captioning and Recognition") as demo:51    gr.Markdown("# ๐Ÿ–ผ๏ธ Image Captioning & Classification App")52    gr.Markdown("Upload an image, then click below to generate a caption or classify it.")53 54    image_input = gr.Image(label="Upload Image", type="pil")55    with gr.Row():56        get_caption_btn = gr.Button("๐Ÿ“ Get Caption")57        classify_btn = gr.Button("๐Ÿ” Classify Image")58    caption_output = gr.Textbox(label="Generated Caption")59    classification_output = gr.Label(label="Top 5 Predictions")60 61    get_caption_btn.click(fn=get_caption, inputs=image_input, outputs=caption_output)62    classify_btn.click(fn=classify_image, inputs=image_input, outputs=classification_output)63 64demo.launch()65