import torch
import torch.nn as nn
from torchvision import models, transforms
from PIL import Image
import numpy as np
import cv2
import matplotlib.pyplot as plt

def gercek_cozum_gradcam(img_path, model_path):
    raw_img = cv2.imread(img_path)
    raw_img = cv2.cvtColor(raw_img, cv2.COLOR_BGR2RGB)
    raw_img = cv2.resize(raw_img, (224, 224))

    img_pil = Image.open(img_path).convert('RGB')
    preprocess = transforms.Compose([
        transforms.Resize((224, 224)),
        transforms.ToTensor(),
        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
    ])
    input_tensor = preprocess(img_pil).unsqueeze(0)

    model = models.vgg19(weights=None)
    model.classifier[6] = nn.Linear(4096, 12) 
    model.load_state_dict(torch.load(model_path, map_location='cpu'))
    model.eval()

    target_layer = model.features[36]
    gradients, activations = [], []
    
    def save_grad(grad): gradients.append(grad)
    def save_act(act): activations.append(act)
    
    h_f = target_layer.register_forward_hook(lambda m, i, o: save_act(o))
    h_b = target_layer.register_full_backward_hook(lambda m, gi, go: save_grad(go[0]))
    
    output = model(input_tensor)
    pred_idx = torch.argmax(output).item()
    output[0, pred_idx].backward()
    
    heatmap = torch.mean(activations[0], dim=1).squeeze().detach().numpy()
    heatmap = np.maximum(heatmap, 0)
    if heatmap.max() > 0: heatmap /= heatmap.max()
    
    h_f.remove()
    h_b.remove()

    heatmap_res = cv2.resize(heatmap, (224, 224))
    heatmap_u8 = np.uint8(255 * heatmap_res)
    heatmap_color = cv2.applyColorMap(heatmap_u8, cv2.COLORMAP_JET)
    heatmap_color = cv2.cvtColor(heatmap_color, cv2.COLOR_BGR2RGB)

    superimposed = cv2.addWeighted(raw_img, 0.6, heatmap_color, 0.4, 0)

    plt.figure(figsize=(12, 6))
    plt.subplot(1, 2, 1)
    plt.imshow(raw_img)
    plt.title("Original Microscopic Image")
    plt.axis('off')

    plt.subplot(1, 2, 2)
    plt.imshow(superimposed)
    plt.title(f"Grad-CAM (Predicted Class Index: {pred_idx})")
    plt.axis('off')
    
    plt.show()

gercek_cozum_gradcam('C:/Users/merkepci/Desktop/code/ornek_resim.jpg', 'C:/Users/merkepci/Desktop/code/vgg19_epoch_best1.pth')

