Spaces:
Sleeping
Sleeping
| 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() |