Spaces:
Running on Zero
Running on Zero
| """ | |
| 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: | |
| 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() |