""" app.py — BrainScan AI (Gradio Space, standalone) ================================================= Aplikasi Gradio mandiri untuk model Hybrid EfficientNet-B3 + Custom ViT dari repo: Marksnb/brain-hybrid-efficientnet-vit Alur: 1. Download checkpoint (.pth) dari Hugging Face Hub saat startup 2. Definisikan arsitektur model (identik dengan classifier_model.py asli) 3. Preprocessing gambar sama seperti saat training (Resize 224 + ImageNet norm) 4. Inference -> probabilitas 5 kelas penyakit otak 5. Generate attention heatmap (ViT attention block terakhir) sebagai visualisasi "area yang difokuskan model" (Explainable AI ringan) Jalankan lokal: pip install gradio torch torchvision huggingface_hub pillow numpy matplotlib python app.py Deploy ke HF Space: - README.md di root Space set: sdk: gradio, app_file: app.py - requirements.txt berisi paket di atas """ import os import numpy as np import torch import torch.nn as nn import torch.nn.functional as F import torchvision.transforms as T from PIL import Image import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import gradio as gr from huggingface_hub import hf_hub_download # ───────────────────────────────────────────────────────────── # 1. KONFIGURASI # ───────────────────────────────────────────────────────────── HF_REPO_ID = "Marksnb/brain-hybrid-efficientnet-vit" CHECKPOINT_FILENAME = "hybrid_vit_efficientnet_brain_best.pth" IMG_SIZE = 224 NUM_CLASSES = 5 CLASSES = [ "Alzheimer", "Intracranial_Hemorrhage", "Normal", "Stroke_Iskemik", "Tumor", ] CLASS_DISPLAY = { "Alzheimer": "Alzheimer", "Intracranial_Hemorrhage": "Intracranial Hemorrhage (ICH)", "Normal": "Normal", "Stroke_Iskemik": "Ischemic Stroke", "Tumor": "Brain Tumor", } IMAGENET_MEAN = [0.485, 0.456, 0.406] IMAGENET_STD = [0.229, 0.224, 0.225] val_transforms = T.Compose([ T.Resize((IMG_SIZE, IMG_SIZE)), T.ToTensor(), T.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD), ]) try: import spaces # noqa: F401 (cek awal, detail di bawah) IS_ZEROGPU = True except ImportError: IS_ZEROGPU = False # Di ZeroGPU Space: GPU baru "muncul" saat fungsi ber-@spaces.GPU dipanggil, # jadi startup HARUS di CPU dulu. Pindah ke cuda dilakukan per-request. if IS_ZEROGPU: DEVICE = torch.device("cpu") RUNTIME_DEVICE = torch.device("cuda") else: DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") RUNTIME_DEVICE = DEVICE # ───────────────────────────────────────────────────────────── # 2. ARSITEKTUR MODEL # (persis sama dengan classifier_model.py di repo Space asli, # supaya checkpoint bisa di-load tanpa error missing/unexpected key) # ───────────────────────────────────────────────────────────── try: from torchvision.models import efficientnet_b3, EfficientNet_B3_Weights HAS_WEIGHTS = True except ImportError: from torchvision.models import efficientnet_b3 HAS_WEIGHTS = False class PatchEmbedding(nn.Module): def __init__(self, in_channels=1536, patch_size=1, embed_dim=768): super().__init__() self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=patch_size) def forward(self, x): x = self.proj(x) x = x.flatten(2).transpose(1, 2) return x class MultiHeadSelfAttention(nn.Module): def __init__(self, embed_dim=768, num_heads=12, dropout=0.1): super().__init__() assert embed_dim % num_heads == 0 self.num_heads = num_heads self.head_dim = embed_dim // num_heads self.scale = self.head_dim ** -0.5 self.qkv = nn.Linear(embed_dim, embed_dim * 3) self.proj = nn.Linear(embed_dim, embed_dim) self.drop = nn.Dropout(dropout) def forward(self, x, return_attn: bool = False): B, N, C = x.shape qkv = (self.qkv(x) .reshape(B, N, 3, self.num_heads, self.head_dim) .permute(2, 0, 3, 1, 4)) q, k, v = qkv[0], qkv[1], qkv[2] attn = (q @ k.transpose(-2, -1)) * self.scale attn = attn.softmax(dim=-1) attn = self.drop(attn) x = (attn @ v).transpose(1, 2).reshape(B, N, C) x = self.proj(x) if return_attn: return x, attn return x class TransformerBlock(nn.Module): def __init__(self, embed_dim=768, num_heads=12, mlp_ratio=4.0, dropout=0.1): super().__init__() self.norm1 = nn.LayerNorm(embed_dim) self.attn = MultiHeadSelfAttention(embed_dim, num_heads, dropout) self.norm2 = nn.LayerNorm(embed_dim) hidden = int(embed_dim * mlp_ratio) self.mlp = nn.Sequential( nn.Linear(embed_dim, hidden), nn.GELU(), nn.Dropout(dropout), nn.Linear(hidden, embed_dim), nn.Dropout(dropout), ) def forward(self, x, return_attn: bool = False): if return_attn: attn_out, attn_weights = self.attn(self.norm1(x), return_attn=True) x = x + attn_out x = x + self.mlp(self.norm2(x)) return x, attn_weights x = x + self.attn(self.norm1(x)) x = x + self.mlp(self.norm2(x)) return x class CrossModalAttentionFusion(nn.Module): def __init__(self, cnn_dim=1536, vit_dim=768, fusion_dim=512, dropout=0.3): super().__init__() self.cnn_proj = nn.Linear(cnn_dim, fusion_dim) self.vit_proj = nn.Linear(vit_dim, fusion_dim) self.attn = nn.Sequential( nn.Linear(fusion_dim * 2, fusion_dim), nn.ReLU(), nn.Linear(fusion_dim, 2), nn.Softmax(dim=-1), ) self.norm = nn.LayerNorm(fusion_dim) self.drop = nn.Dropout(dropout) def forward(self, cnn_feat, vit_feat): c = self.cnn_proj(cnn_feat) v = self.vit_proj(vit_feat) w = self.attn(torch.cat([c, v], dim=-1)) fused = w[:, 0:1] * c + w[:, 1:2] * v fused = self.norm(fused) fused = self.drop(fused) return fused class BrainHybridModel(nn.Module): def __init__(self, num_classes: int = NUM_CLASSES, vit_embed_dim: int = 768, vit_num_heads: int = 12, vit_num_layers: int = 6, fusion_dim: int = 512, dropout: float = 0.3, freeze_backbone: bool = True): super().__init__() if HAS_WEIGHTS: backbone = efficientnet_b3(weights=EfficientNet_B3_Weights.DEFAULT) else: backbone = efficientnet_b3(pretrained=True) self.features = backbone.features self.cnn_out = 1536 self.patch_embed = PatchEmbedding(self.cnn_out, patch_size=1, embed_dim=vit_embed_dim) self.cls_token = nn.Parameter(torch.zeros(1, 1, vit_embed_dim)) nn.init.trunc_normal_(self.cls_token, std=0.02) num_patches = (IMG_SIZE // 32) ** 2 self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, vit_embed_dim)) nn.init.trunc_normal_(self.pos_embed, std=0.02) self.pos_drop = nn.Dropout(dropout) self.blocks = nn.ModuleList([ TransformerBlock(vit_embed_dim, vit_num_heads, dropout=dropout) for _ in range(vit_num_layers) ]) self.vit_norm = nn.LayerNorm(vit_embed_dim) self.fusion = CrossModalAttentionFusion( cnn_dim=self.cnn_out, vit_dim=vit_embed_dim, fusion_dim=fusion_dim, dropout=dropout) self.classifier = nn.Sequential( nn.Linear(fusion_dim, 256), nn.GELU(), nn.BatchNorm1d(256), nn.Dropout(dropout), nn.Linear(256, num_classes), ) if freeze_backbone: for param in self.features.parameters(): param.requires_grad = False def forward(self, x): feat_map = self.features(x) cnn_feat = F.adaptive_avg_pool2d(feat_map, 1).flatten(1) patches = self.patch_embed(feat_map) cls = self.cls_token.expand(x.size(0), -1, -1) tokens = torch.cat([cls, patches], dim=1) tokens = tokens + self.pos_embed tokens = self.pos_drop(tokens) for blk in self.blocks: tokens = blk(tokens) tokens = self.vit_norm(tokens) vit_feat = tokens[:, 0] fused = self.fusion(cnn_feat, vit_feat) logits = self.classifier(fused) return logits def forward_with_attention(self, x): feat_map = self.features(x) cnn_feat = F.adaptive_avg_pool2d(feat_map, 1).flatten(1) patches = self.patch_embed(feat_map) cls = self.cls_token.expand(x.size(0), -1, -1) tokens = torch.cat([cls, patches], dim=1) tokens = tokens + self.pos_embed tokens = self.pos_drop(tokens) last_attn = None for i, blk in enumerate(self.blocks): if i == len(self.blocks) - 1: tokens, last_attn = blk(tokens, return_attn=True) else: tokens = blk(tokens) tokens = self.vit_norm(tokens) vit_feat = tokens[:, 0] fused = self.fusion(cnn_feat, vit_feat) logits = self.classifier(fused) return logits, last_attn # ───────────────────────────────────────────────────────────── # 3. LOAD MODEL (sekali saat startup) # ───────────────────────────────────────────────────────────── print(f"[startup] Downloading checkpoint '{CHECKPOINT_FILENAME}' dari {HF_REPO_ID} ...") checkpoint_path = hf_hub_download(repo_id=HF_REPO_ID, filename=CHECKPOINT_FILENAME) print(f"[startup] Checkpoint tersimpan di: {checkpoint_path}") model = BrainHybridModel().to(DEVICE) state_dict = torch.load(checkpoint_path, map_location=DEVICE) # Beberapa checkpoint training disimpan sebagai dict {"model_state_dict": ...} if isinstance(state_dict, dict) and "model_state_dict" in state_dict: state_dict = state_dict["model_state_dict"] missing, unexpected = model.load_state_dict(state_dict, strict=False) if missing: print(f"[startup] WARNING - missing keys: {missing}") if unexpected: print(f"[startup] WARNING - unexpected keys: {unexpected}") model.eval() print(f"[startup] Model siap. Device: {DEVICE}") # ───────────────────────────────────────────────────────────── # 4. FUNGSI INFERENCE + ATTENTION HEATMAP # ───────────────────────────────────────────────────────────── def generate_attention_overlay(orig_image: Image.Image, tensor_image: torch.Tensor, attn: torch.Tensor): """Buat gambar overlay heatmap attention (ViT) di atas gambar asli.""" avg_attn = attn.squeeze(0).mean(dim=0) # [seq_len, seq_len] cls_attn = avg_attn[0, 1:] # attention CLS -> semua patch num_patches = int(cls_attn.shape[0] ** 0.5) heatmap = cls_attn.reshape(num_patches, num_patches).cpu().numpy() heatmap = np.maximum(heatmap, 0) heatmap = heatmap / (np.max(heatmap) if np.max(heatmap) != 0 else 1.0) heatmap_img = Image.fromarray((heatmap * 255).astype(np.uint8)) heatmap_resized = np.array( heatmap_img.resize(orig_image.size, Image.Resampling.BILINEAR) ) / 255.0 fig, ax = plt.subplots(figsize=(5, 5)) ax.imshow(orig_image) ax.imshow(heatmap_resized, cmap="jet", alpha=0.45) ax.axis("off") ax.set_title("Peta Fokus Atensi AI (ViT Attention)") fig.tight_layout() fig.canvas.draw() overlay_img = Image.frombytes("RGB", fig.canvas.get_width_height(), fig.canvas.tostring_rgb()) plt.close(fig) return overlay_img def _analyze_brain_scan_impl(image: Image.Image): if image is None: return None, None, "Silakan upload gambar CT-Scan / MRI otak terlebih dahulu." infer_device = RUNTIME_DEVICE if IS_ZEROGPU else DEVICE model.to(infer_device) orig_image = image.convert("RGB") tensor_image = val_transforms(orig_image).unsqueeze(0).to(infer_device) with torch.no_grad(): logits, attn = model.forward_with_attention(tensor_image) probs = F.softmax(logits, dim=1).squeeze(0).cpu().numpy() pred_idx = int(np.argmax(probs)) pred_class = CLASSES[pred_idx] pred_label = CLASS_DISPLAY[pred_class] confidence = float(probs[pred_idx]) * 100 # Dict untuk gr.Label (semua kelas + probabilitasnya) label_scores = {CLASS_DISPLAY[c]: float(p) for c, p in zip(CLASSES, probs)} overlay_img = generate_attention_overlay(orig_image, tensor_image, attn) summary = ( f"**Prediksi: {pred_label}** (keyakinan {confidence:.2f}%)\n\n" f"Catatan: hasil ini adalah output model AI, BUKAN diagnosis medis resmi. " f"Selalu konsultasikan dengan dokter/radiolog untuk keputusan klinis." ) return label_scores, overlay_img, summary if IS_ZEROGPU: @spaces.GPU def analyze_brain_scan(image: Image.Image): return _analyze_brain_scan_impl(image) else: def analyze_brain_scan(image: Image.Image): return _analyze_brain_scan_impl(image) # ───────────────────────────────────────────────────────────── # 5. UI GRADIO # ───────────────────────────────────────────────────────────── with gr.Blocks(title="BrainScan AI — Hybrid EfficientNet-ViT") as demo: gr.Markdown( """ # 🧠 BrainScan AI Klasifikasi otomatis CT-Scan / MRI otak menggunakan arsitektur **Hybrid EfficientNet-B3 + Custom Vision Transformer** dengan Cross-Modal Attention Fusion. Kelas yang dideteksi: Alzheimer, Intracranial Hemorrhage (ICH), Normal, Ischemic Stroke, Brain Tumor. ⚠️ **Disclaimer:** alat ini untuk tujuan riset/edukasi, bukan pengganti diagnosis medis profesional. """ ) with gr.Row(): with gr.Column(): image_input = gr.Image(type="pil", label="Upload CT-Scan / MRI Otak") analyze_btn = gr.Button("🔍 Analisis", variant="primary") with gr.Column(): label_output = gr.Label(num_top_classes=5, label="Probabilitas per Kelas") heatmap_output = gr.Image(label="Peta Fokus Atensi AI (Explainability)") summary_output = gr.Markdown() analyze_btn.click( fn=analyze_brain_scan, inputs=image_input, outputs=[label_output, heatmap_output, summary_output], api_name="analyze", ) gr.Examples( examples=[], # tambahkan path gambar contoh di sini kalau ada, mis. "samples/normal_1.jpg" inputs=image_input, ) if __name__ == "__main__": demo.launch()