brainscan / app.py
Marksnb's picture
Upload app.py
3c206da verified
Raw
History Blame Contribute Delete
23.5 kB
"""
app.py β€” BrainScan AI (Gradio Space, standalone, single-file)
================================================================
Model: Hybrid EfficientNet-B3 + Custom ViT (Cross-Modal Attention Fusion)
Repo checkpoint : Marksnb/brain-hybrid-efficientnet-vit
- hybrid_vit_efficientnet_brain_best.pth -> model klasifikasi 5 kelas (utama)
- best_precheck_model.pth -> model precheck biner
("apakah gambar ini CT/MRI otak?")
Alur:
1. Download kedua checkpoint dari Hugging Face Hub saat startup.
2. Definisikan arsitektur model utama (Hybrid EfficientNet-B3 + ViT).
3. Load model precheck secara ADAPTIF: mencoba beberapa arsitektur backbone
kandidat dan memilih yang paling cocok dengan checkpoint (lihat catatan
di bagian PRECHECK MODEL di bawah -- arsitektur aslinya tidak
didokumentasikan di repo, jadi ini best-effort & auto-degrade jika
tidak cocok).
4. Preprocessing gambar sama seperti saat training (Resize 224 + ImageNet norm).
5. Inference -> precheck dulu, baru klasifikasi 5 kelas penyakit otak.
6. Generate attention heatmap (ViT attention block terakhir) sebagai
visualisasi "area yang difokuskan model" (Explainable AI ringan).
Jalankan lokal:
pip install -r requirements.txt
python app.py
Deploy ke HF Space:
- README.md di root Space (metadata YAML): sdk: gradio, app_file: app.py
- requirements.txt berisi paket yang dibutuhkan
- Endpoint REST otomatis tersedia di /gradio_api/call/analyze
(lihat api_name="analyze" di bagian UI paling bawah)
"""
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"
MAIN_CHECKPOINT_FILENAME = "hybrid_vit_efficientnet_brain_best.pth"
PRECHECK_CHECKPOINT_FILENAME = "best_precheck_model.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),
])
# --- Precheck config -------------------------------------------------------
# CATATAN PENTING: arsitektur & urutan kelas model precheck TIDAK
# didokumentasikan di repo HF, jadi ini asumsi. Index 0 = bukan brain scan,
# index 1 = brain scan. Kalau hasil precheck kebalik-balik setelah deploy,
# tinggal tukar dua string ini.
PRECHECK_CLASS_NAMES = ["Bukan_Brain_Scan", "Brain_Scan"]
# Ambang keyakinan minimum supaya precheck menolak gambar (0-1).
PRECHECK_REJECT_THRESHOLD = 0.65
try:
import spaces # noqa: F401
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 UTAMA (Hybrid EfficientNet-B3 + Custom ViT)
# (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_ENUM = True
except ImportError:
from torchvision.models import efficientnet_b3
HAS_WEIGHTS_ENUM = 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,
pretrained_backbone: bool = True):
super().__init__()
if pretrained_backbone:
if HAS_WEIGHTS_ENUM:
backbone = efficientnet_b3(weights=EfficientNet_B3_Weights.DEFAULT)
else:
backbone = efficientnet_b3(pretrained=True)
else:
backbone = efficientnet_b3(weights=None) if HAS_WEIGHTS_ENUM else efficientnet_b3(pretrained=False)
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 _encode(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)
return cnn_feat, tokens
def forward(self, x):
cnn_feat, tokens = self._encode(x)
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):
cnn_feat, tokens = self._encode(x)
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
def _unwrap_state_dict(raw):
"""Beberapa checkpoint disimpan sebagai dict {'model_state_dict': ...} atau
{'state_dict': ...}. Fungsi ini menormalkannya menjadi state_dict polos,
dan membuang prefix 'module.' (umum kalau training pakai DataParallel)."""
if isinstance(raw, dict):
for key in ("model_state_dict", "state_dict", "model"):
if key in raw and isinstance(raw[key], dict):
raw = raw[key]
break
cleaned = {}
for k, v in raw.items():
cleaned[k.replace("module.", "", 1) if k.startswith("module.") else k] = v
return cleaned
# ─────────────────────────────────────────────────────────────
# 3. MODEL PRECHECK (biner: brain scan vs bukan)
# ARSITEKTUR ASLINYA TIDAK DIDOKUMENTASIKAN DI REPO -> kita coba beberapa
# backbone ringan yang umum dipakai untuk precheck/gatekeeper model, dan
# pilih otomatis yang paling cocok (paling sedikit missing/unexpected key)
# dengan checkpoint. Kalau tidak ada yang cukup cocok, precheck otomatis
# dimatikan (app tetap jalan, hanya tanpa langkah precheck).
# ─────────────────────────────────────────────────────────────
def _build_precheck_candidates(num_out=2):
"""Kembalikan list (nama, model) kandidat arsitektur backbone ringan
dengan output akhir num_out kelas."""
import torchvision.models as tvm
candidates = []
def safe(name, fn):
try:
candidates.append((name, fn()))
except Exception as e:
print(f"[precheck] Lewati kandidat '{name}': {e}")
def make_resnet18():
m = tvm.resnet18(weights=None)
m.fc = nn.Linear(m.fc.in_features, num_out)
return m
def make_resnet34():
m = tvm.resnet34(weights=None)
m.fc = nn.Linear(m.fc.in_features, num_out)
return m
def make_mobilenet_v2():
m = tvm.mobilenet_v2(weights=None)
m.classifier[-1] = nn.Linear(m.classifier[-1].in_features, num_out)
return m
def make_efficientnet_b0():
m = tvm.efficientnet_b0(weights=None)
m.classifier[-1] = nn.Linear(m.classifier[-1].in_features, num_out)
return m
def make_densenet121():
m = tvm.densenet121(weights=None)
m.classifier = nn.Linear(m.classifier.in_features, num_out)
return m
safe("resnet18", make_resnet18)
safe("resnet34", make_resnet34)
safe("mobilenet_v2", make_mobilenet_v2)
safe("efficientnet_b0", make_efficientnet_b0)
safe("densenet121", make_densenet121)
return candidates
def load_precheck_model(checkpoint_path, device):
"""Coba beberapa arsitektur kandidat, pilih yang paling cocok dengan
checkpoint. Return (model_or_None, info_string)."""
try:
raw = torch.load(checkpoint_path, map_location="cpu")
except Exception as e:
return None, f"Gagal membaca checkpoint precheck: {e}"
state_dict = _unwrap_state_dict(raw)
total_keys = max(len(state_dict), 1)
best = None # (score, name, model)
for name, model in _build_precheck_candidates():
model_keys = set(model.state_dict().keys())
missing, unexpected = model.load_state_dict(state_dict, strict=False)
# `missing`/`unexpected` di sini adalah namedtuple hasil load_state_dict
n_bad = len(missing) + len(unexpected)
score = 1.0 - (n_bad / total_keys)
print(f"[precheck] Kandidat '{name}': score={score:.3f} "
f"(missing={len(missing)}, unexpected={len(unexpected)})")
if best is None or score > best[0]:
best = (score, name, model)
if best is None:
return None, "Tidak ada kandidat arsitektur yang bisa dibangun."
score, name, model = best
if score < 0.9:
return None, (f"Precheck dinonaktifkan: arsitektur checkpoint tidak "
f"cocok dengan kandidat manapun (skor terbaik={score:.2f}, "
f"kandidat={name}). Cek log server untuk detail key yang "
f"tidak cocok, lalu sesuaikan _build_precheck_candidates().")
model.to(device)
model.eval()
return model, f"Precheck aktif menggunakan arsitektur '{name}' (skor kecocokan={score:.2f})."
# ─────────────────────────────────────────────────────────────
# 4. LOAD MODEL (sekali saat startup)
# ─────────────────────────────────────────────────────────────
print(f"[startup] Downloading '{MAIN_CHECKPOINT_FILENAME}' dari {HF_REPO_ID} ...")
main_checkpoint_path = hf_hub_download(repo_id=HF_REPO_ID, filename=MAIN_CHECKPOINT_FILENAME)
print(f"[startup] Checkpoint utama tersimpan di: {main_checkpoint_path}")
print(f"[startup] Downloading '{PRECHECK_CHECKPOINT_FILENAME}' dari {HF_REPO_ID} ...")
precheck_checkpoint_path = hf_hub_download(repo_id=HF_REPO_ID, filename=PRECHECK_CHECKPOINT_FILENAME)
print(f"[startup] Checkpoint precheck tersimpan di: {precheck_checkpoint_path}")
model = BrainHybridModel().to(DEVICE)
raw_state = torch.load(main_checkpoint_path, map_location=DEVICE)
main_state_dict = _unwrap_state_dict(raw_state)
missing, unexpected = model.load_state_dict(main_state_dict, strict=False)
if missing:
print(f"[startup] WARNING - model utama, missing keys: {missing}")
if unexpected:
print(f"[startup] WARNING - model utama, unexpected keys: {unexpected}")
model.eval()
print(f"[startup] Model utama siap. Device: {DEVICE}")
precheck_model, precheck_info = load_precheck_model(precheck_checkpoint_path, DEVICE)
PRECHECK_ENABLED = precheck_model is not None
print(f"[startup] {precheck_info}")
# ─────────────────────────────────────────────────────────────
# 5. FUNGSI INFERENCE + ATTENTION HEATMAP
# ─────────────────────────────────────────────────────────────
def generate_attention_overlay(orig_image: Image.Image, 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).detach().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()
# buffer_rgba() kompatibel dengan matplotlib versi baru (tostring_rgb
# sudah deprecated/dihapus di beberapa versi terbaru).
buf = np.asarray(fig.canvas.buffer_rgba())
overlay_img = Image.fromarray(buf).convert("RGB")
plt.close(fig)
return overlay_img
def _run_precheck(tensor_image: torch.Tensor):
"""Return (is_brain_scan: bool, confidence: float, label_scores: dict)."""
with torch.no_grad():
logits = precheck_model(tensor_image)
probs = F.softmax(logits, dim=1).squeeze(0).cpu().numpy()
pred_idx = int(np.argmax(probs))
label_scores = {PRECHECK_CLASS_NAMES[i]: float(p) for i, p in enumerate(probs)}
is_brain_scan = (pred_idx == 1)
confidence = float(probs[pred_idx])
return is_brain_scan, confidence, label_scores
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)
if PRECHECK_ENABLED:
precheck_model.to(infer_device)
orig_image = image.convert("RGB")
tensor_image = val_transforms(orig_image).unsqueeze(0).to(infer_device)
precheck_note = ""
if PRECHECK_ENABLED:
is_brain_scan, pc_conf, _ = _run_precheck(tensor_image)
if (not is_brain_scan) and pc_conf >= PRECHECK_REJECT_THRESHOLD:
warning = (
f"⚠️ **Gambar ini kemungkinan BUKAN CT-Scan/MRI otak** "
f"(keyakinan precheck {pc_conf * 100:.1f}%).\n\n"
f"Model klasifikasi utama tidak dijalankan karena gambar tidak "
f"lolos precheck. Silakan upload ulang dengan gambar CT-Scan "
f"atau MRI otak yang valid."
)
return None, None, warning
precheck_note = f"βœ… Precheck: gambar terdeteksi sebagai brain scan (keyakinan {pc_conf * 100:.1f}%).\n\n"
else:
precheck_note = "ℹ️ Precheck dinonaktifkan (lihat log server untuk detail).\n\n"
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
label_scores = {CLASS_DISPLAY[c]: float(p) for c, p in zip(CLASSES, probs)}
overlay_img = generate_attention_overlay(orig_image, attn)
summary = (
f"{precheck_note}"
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(duration=60)
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)
# ─────────────────────────────────────────────────────────────
# 6. 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, dilengkapi model **precheck** untuk
memvalidasi apakah gambar yang diupload benar-benar CT/MRI otak.
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()