brain / explainability.py
Marksnb's picture
Upload 13 files
12d60da verified
Raw
History Blame Contribute Delete
3.49 kB
import os
import torch
import numpy as np
import matplotlib.pyplot as plt
from PIL import Image
from src.preprocess import val_transforms
from src.models.classifier_model import BrainHybridModel
from src.config import OUTPUT_DIR, CHECKPOINT_DIR
def generate_attention_heatmap(image_path, save_name="attention_map.png"):
"""Menghasilkan peta panas (heatmap) fokus perhatian model AI pada gambar otak"""
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# 1. Muat dan ubah gambar menjadi tensor
orig_image = Image.open(image_path).convert('RGB')
tensor_image = val_transforms(orig_image).unsqueeze(0).to(device)
# 2. Muat model dan bobot terbaik
model = BrainHybridModel().to(device)
checkpoint_path = os.path.join(CHECKPOINT_DIR, "best_hybrid_model.pth")
if os.path.exists(checkpoint_path):
try:
model.load_state_dict(torch.load(checkpoint_path, map_location=device))
except RuntimeError:
print("Warning: Checkpoint tidak kompatibel dengan arsitektur ViT baru, menggunakan bobot pretrained bawaan.")
model.eval()
# 3. Ekstraksi attention weights dari ViT (custom Transformer block terakhir)
with torch.no_grad():
# forward_with_attention mengembalikan (logits, attn) dari block terakhir
# attn shape: [B, num_heads, seq_len, seq_len]
_, attentions = model.forward_with_attention(tensor_image)
# Rata-ratakan semua attention heads
avg_attn = attentions.squeeze(0).mean(dim=0) # [seq_len, seq_len]
# Ambil attention dari CLS token (index 0) ke semua patch tokens
cls_attn = avg_attn[0, 1:] # [num_patches] (buang CLS-to-CLS)
# Feature map EfficientNet-B3 di 224x224 -> grid 7x7 = 49 patch tokens
num_patches = int(cls_attn.shape[0] ** 0.5)
heatmap = cls_attn.reshape(num_patches, num_patches).cpu().numpy()
# Normalisasi peta panas antara nilai 0 hingga 1
heatmap = np.maximum(heatmap, 0)
heatmap /= np.max(heatmap) if np.max(heatmap) != 0 else 1.0
# 4. Gambar dan gabungkan citra asli dengan peta panas
fig, axes = plt.subplots(1, 2, figsize=(10, 5))
axes[0].imshow(orig_image)
axes[0].set_title("Gambar Medis Asli")
axes[0].axis('off')
# Ubah ukuran peta panas agar pas dengan dimensi gambar asli
heatmap_resized = np.array(Image.fromarray(heatmap).resize(orig_image.size, Image.Resampling.BILINEAR))
axes[1].imshow(orig_image)
axes[1].imshow(heatmap_resized, cmap='jet', alpha=0.4) # Overlay warna transparan
axes[1].set_title("Peta Fokus Atensi AI (ViT Attention)")
axes[1].axis('off')
# Simpan visualisasi ke folder outputs/figures/
figure_dir = os.path.join(OUTPUT_DIR, "figures")
os.makedirs(figure_dir, exist_ok=True)
save_path = os.path.join(figure_dir, save_name)
plt.savefig(save_path, bbox_inches='tight')
plt.close()
print(f"Sukses menghasilkan peta eksplanabilitas AI! Tersimpan di: {save_path}")
if __name__ == "__main__":
# Mencari satu contoh gambar acak dari folder normal untuk uji coba modul
sample_dir = "data/raw/Normal"
if os.path.exists(sample_dir) and os.listdir(sample_dir):
first_img = os.listdir(sample_dir)[0]
full_path = os.path.join(sample_dir, first_img)
generate_attention_heatmap(full_path)
else:
print("Folder data/raw/Normal kosong atau tidak ditemukan untuk pengujian.")