import gradio as gr import torch from torchvision import models, transforms from PIL import Image from grad_cam import compute_heatmap, upsampleHeatmap # model model = models.resnet18(pretrained=True) model.eval() # transform image_transform = transforms.Compose([ transforms.Resize((224,224)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) def predict(image): # image = PIL image from Gradio image_tensor = image_transform(image).unsqueeze(0) heatmap, pred_id = compute_heatmap(model, image_tensor) overlay, _ = upsampleHeatmap(heatmap, image_tensor) with open("imagenet_classes.txt", "r") as f: labels = f.read().splitlines() pred_label = labels[pred_id] return overlay, f"Predicted class: {pred_label}" demo = gr.Interface( fn=predict, inputs=gr.Image(type="pil"), outputs=[ gr.Image(type="numpy", label="Grad-CAM"), gr.Text(label="Prediction") ], title="Grad-CAM Explainability Demo" ) demo.launch()