ResNet50-P3 / gradcam_visualizer.py
chrisnguyenx's picture
Upload gradcam_visualizer.py with huggingface_hub
35356ce verified
Raw
History Blame Contribute Delete
5.85 kB
"""
gradcam_visualizer.py
====================
Script sinh bản đồ nhiệt Grad-CAM (Gradient-weighted Class Activation Mapping)
giải thích vùng chú ý y khoa của mô hình ResNet-50.
Sử dụng:
python gradcam_visualizer.py --image path/to/fundus.png --output path/to/gradcam_result.png
"""
from __future__ import annotations
import os
import sys
import argparse
import numpy as np
import cv2
import torch
import torch.nn.functional as F
from typing import Union
from PIL import Image
try:
from .predictor import DRPredictor
from .preprocessing import prepare_image_tensor, full_preprocess_pipeline, load_image
except ImportError:
from predictor import DRPredictor
from preprocessing import prepare_image_tensor, full_preprocess_pipeline, load_image
class GradCAM:
"""
Lớp Grad-CAM chuyên biệt cho mô hình ResNet-50.
Trích xuất gradient và activation từ lớp convolution cuối cùng (layer4[-1]).
"""
def __init__(self, model: torch.nn.Module, target_layer: torch.nn.Module):
self.model = model
self.target_layer = target_layer
self.gradients = None
self.activations = None
# Register forward & backward hooks
self.target_layer.register_forward_hook(self._save_activation)
self.target_layer.register_full_backward_hook(self._save_gradient)
def _save_activation(self, module, input, output):
self.activations = output.detach()
def _save_gradient(self, module, grad_input, grad_output):
self.gradients = grad_output[0].detach()
def generate(self, input_tensor: torch.Tensor, target_class: int | None = None) -> tuple[np.ndarray, int]:
"""
Sinh heatmap Grad-CAM cho 1 tensor ảnh đầu vào (1, C, H, W).
"""
self.model.eval()
self.model.zero_grad()
output = self.model(input_tensor)
if target_class is None:
target_class = torch.argmax(output, dim=1).item()
score = output[0, target_class]
score.backward()
gradients = self.gradients[0] # (C, H_feat, W_feat)
activations = self.activations[0] # (C, H_feat, W_feat)
# Global Average Pooling trên gradients để tính trọng số alpha_k
weights = torch.mean(gradients, dim=(1, 2), keepdim=True) # (C, 1, 1)
# Tổng có trọng số của các kênh activation
cam = torch.sum(weights * activations, dim=0) # (H_feat, W_feat)
# Áp dụng ReLU
cam = F.relu(cam)
# Chuẩn hóa về [0, 1]
cam = cam.cpu().numpy()
if cam.max() > 0:
cam = cam / cam.max()
return cam, target_class
def overlay_cam_on_image(
img_bgr: np.ndarray,
cam: np.ndarray,
alpha: float = 0.5,
colormap: int = cv2.COLORMAP_JET,
) -> np.ndarray:
"""
Chồng heatmap Grad-CAM màu JET lên ảnh BGR gốc.
"""
h, w = img_bgr.shape[:2]
cam_resized = cv2.resize(cam, (w, h))
heatmap = cv2.applyColorMap(np.uint8(255 * cam_resized), colormap)
blended = cv2.addWeighted(img_bgr, 1 - alpha, heatmap, alpha, 0)
return blended
def generate_gradcam_for_image(
predictor: DRPredictor,
image_input: Union[str, bytes, Image.Image, np.ndarray],
output_path: str | None = None,
target_class: int | None = None,
draw_label: bool = False,
) -> np.ndarray:
"""
Hàm wrapper hoàn chỉnh: Đọc ảnh -> Tiền xử lý -> Chạy Grad-CAM -> Lưu/trả về ảnh BGR kết quả.
"""
if isinstance(image_input, str):
if not os.path.exists(image_input):
raise FileNotFoundError(f"Không tìm thấy ảnh: {image_input}")
img_raw = cv2.imread(image_input)
else:
img_raw = load_image(image_input)
img_preprocessed = full_preprocess_pipeline(
img_raw, target_size=predictor.img_size, use_ben_graham=True
)
tensor_img = prepare_image_tensor(
image_input=img_raw,
target_size=predictor.img_size,
mean=predictor.mean,
std=predictor.std,
use_ben_graham=True,
).unsqueeze(0).to(predictor.device)
# Attach GradCAM hook onto the last ResNet block (model.layer4[-1])
target_layer = predictor.model.model.layer4[-1]
grad_cam = GradCAM(predictor.model, target_layer)
cam, pred_class = grad_cam.generate(tensor_img, target_class=target_class)
class_name = predictor.class_names.get(str(pred_class), f"Class {pred_class}")
blended = overlay_cam_on_image(img_preprocessed, cam, alpha=0.5)
# Nếu truyền draw_label=True thì mới vẽ chữ lên ảnh
if draw_label:
label_text = f"Pred: {class_name} (Class {pred_class})"
cv2.putText(blended, label_text, (10, 25), cv2.FONT_HERSHEY_SIMPLEX, 0.7, (255, 255, 255), 2)
if output_path:
os.makedirs(os.path.dirname(os.path.abspath(output_path)), exist_ok=True)
cv2.imwrite(output_path, blended)
print(f"[OK] Da tao Grad-CAM heatmap -> {output_path}")
return blended
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Sinh Grad-CAM Heatmap giải thích mô hình DR (ResNet-50)")
parser.add_argument("--image", type=str, required=True, help="Đường dẫn tới ảnh đáy mắt")
parser.add_argument("--output", type=str, default="gradcam_output_resnet.png", help="Đường dẫn lưu ảnh Grad-CAM")
parser.add_argument("--class_id", type=int, default=None, help="Class ID muốn giải thích (nếu None sẽ lấy Argmax)")
parser.add_argument("--draw_label", action="store_true", help="Vẽ nhãn văn bản lên ảnh")
args = parser.parse_args()
predictor = DRPredictor()
generate_gradcam_for_image(predictor, args.image, args.output, target_class=args.class_id, draw_label=args.draw_label)