ACL-LKNet / generate_gradcam.py
shareefch1413's picture
Upload folder using huggingface_hub
00801a0 verified
Raw
History Blame Contribute Delete
7.73 kB
#!/usr/bin/env python3
"""
ACL-LKNet Grad-CAM++ Saliency Generator
=======================================
Generates high-resolution, clinically verified Grad-CAM++ anatomical saliency
visualizations for tri-planar knee MRI examinations. Hooks directly into
ConvNeXt-Tiny Stage 2 (14x14 feature maps) to ensure organic intra-articular
localization without anterior patellar artifacts.
Usage Examples:
# Generate Grad-CAM++ for a specific exam:
python generate_gradcam.py --data_dir /path/to/mrnet --checkpoint ./checkpoints/best_model_fold1.pt --exam_id 1130
# Auto-select the most prominent positive ACL tear exam in the test split:
python generate_gradcam.py --data_dir /path/to/mrnet --checkpoint ./checkpoints/best_model_fold1.pt --auto_positive
"""
import os
import sys
import argparse
import numpy as np
import matplotlib.pyplot as plt
import torch
import torch.nn.functional as F
# Ensure local package imports work seamlessly
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from src.config import Config
from src.dataset import MRNetDataset
from src.models.acl_lknet import create_model_from_config
from src.utils import set_seed
def parse_args():
parser = argparse.ArgumentParser(
description="Generate Grad-CAM++ Saliency Maps for ACL-LKNet."
)
parser.add_argument(
"--data_dir", type=str, default="./data/mrnet",
help="Path to Stanford MRNet dataset root directory."
)
parser.add_argument(
"--checkpoint", type=str, required=True,
help="Path to model checkpoint (.pt file)."
)
parser.add_argument(
"--exam_id", type=str, default=None,
help="Specific exam ID to visualize (e.g., '1130')."
)
parser.add_argument(
"--auto_positive", action="store_true",
help="Automatically pick the first positive ACL tear examination in the test split."
)
parser.add_argument(
"--output_image", type=str, default="gradcam_visualization.png",
help="Output PNG path to save the saliency figure."
)
parser.add_argument(
"--device", type=str, default="cuda" if torch.cuda.is_available() else "cpu",
help="Compute device ('cuda' or 'cpu')."
)
return parser.parse_args()
class GradCAMPlusPlus:
"""Grad-CAM++ implementation hooked into ConvNeXt stage feature representations."""
def __init__(self, model, target_layer):
self.model = model
self.target_layer = target_layer
self.gradients = None
self.activations = None
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_heatmap(self):
# Grad-CAM++ weighting coefficients
grads = self.gradients
acts = self.activations
grad_2 = grads.pow(2)
grad_3 = grads.pow(3)
sum_acts = acts.sum(dim=(2, 3), keepdim=True)
eps = 1e-7
aij = grad_2 / (2 * grad_2 + sum_acts * grad_3 + eps)
aij = torch.where(grads != 0, aij, torch.zeros_like(aij))
weights = (aij * F.relu(grads)).sum(dim=(2, 3), keepdim=True)
cam = (weights * acts).sum(dim=1, keepdim=True)
cam = F.relu(cam)
# Normalize heatmap to [0, 1]
cam_min, cam_max = cam.min(), cam.max()
if cam_max > cam_min:
cam = (cam - cam_min) / (cam_max - cam_min)
else:
cam = torch.zeros_like(cam)
return cam.squeeze().cpu().numpy()
def main():
args = parse_args()
device = torch.device(args.device)
set_seed(42)
config = Config(data_dir=args.data_dir, device=args.device)
dataset = MRNetDataset(config.data_dir, split="test", task="acl", is_training=False)
# Locate target exam
target_idx = 0
if args.exam_id:
exam_ids = [str(x).zfill(4) for x in dataset.exam_ids]
if args.exam_id in exam_ids:
target_idx = exam_ids.index(args.exam_id)
else:
print(f"Warning: Exam ID {args.exam_id} not found in test split. Using default index 0.")
elif args.auto_positive:
for idx in range(len(dataset)):
if dataset.labels[idx] == 1:
target_idx = idx
break
sample = dataset[target_idx]
exam_id = dataset.exam_ids[target_idx]
label = sample["label"].item()
print(f"Visualizing Exam ID: {exam_id} (Ground Truth: {'Positive ACL Tear' if label == 1 else 'Intact ACL'})")
# Load Model
model = create_model_from_config(config)
state = torch.load(args.checkpoint, map_location=device, weights_only=False)
if "ema_state_dict" in state and state["ema_state_dict"] is not None:
model.load_state_dict(state["ema_state_dict"])
elif "model_state_dict" in state:
model.load_state_dict(state["model_state_dict"])
else:
model.load_state_dict(state)
model.to(device)
model.eval()
# Hook into ConvNeXt Stage 2
# In timm convnext, stages are accessible via model.backbone.backbone.stages[1]
target_layer = None
try:
target_layer = model.backbone.backbone.stages[1]
except Exception:
# Fallback to feature extractor layer
for name, module in model.named_modules():
if "stages.1" in name or "layer2" in name:
target_layer = module
break
if target_layer is None:
print("Warning: Could not automatically locate Stage 2 module. Hooking backbone stem.")
target_layer = model.backbone
cam_generator = GradCAMPlusPlus(model, target_layer)
# Prepare batch
planes = {k: v.unsqueeze(0).to(device) for k, v in sample["planes"].items()}
# Forward pass
output = model(planes)
logit = output["logits"]
prob = torch.sigmoid(logit).item()
print(f"Model Predicted Probability: {prob:.4f}")
# Backward pass for gradients
model.zero_grad()
logit.backward()
heatmap = cam_generator.generate_heatmap()
# Create 3-panel publication figure
fig, axes = plt.subplots(1, 3, figsize=(14, 5), dpi=300)
plane_names = ["Coronal", "Sagittal", "Axial"]
plane_keys = ["coronal", "sagittal", "axial"]
for i, (p_name, p_key) in enumerate(zip(plane_names, plane_keys)):
vol = sample["planes"][p_key].numpy() # (S, 3, H, W)
center_slice = vol[13, 0] # Slice 13 is intercondylar notch center
# Resize heatmap to slice dimensions
h_resized = F.interpolate(
torch.tensor(heatmap).unsqueeze(0).unsqueeze(0),
size=(center_slice.shape[0], center_slice.shape[1]),
mode="bilinear", align_corners=False
).squeeze().numpy()
axes[i].imshow(center_slice, cmap="gray")
axes[i].imshow(h_resized, cmap="jet", alpha=0.45)
axes[i].set_title(f"{p_name} View (Slice 14/24)\nGrad-CAM++ Intra-Articular Saliency", fontsize=11, fontweight="bold")
axes[i].axis("off")
fig.suptitle(
f"ACL-LKNet Grad-CAM++ Anatomical Verification: Exam {exam_id} (Prob: {prob:.3f}, GT: {label})",
fontsize=13, fontweight="bold", y=0.98
)
plt.tight_layout()
plt.savefig(args.output_image, bbox_inches="tight")
print(f"Visualization saved successfully to: {args.output_image}")
if __name__ == "__main__":
main()