Pneumonia-Detection / src /streamlit_app.py
Spitblaze's picture
Update src/streamlit_app.py
60a960e verified
Raw
History Blame Contribute Delete
2.42 kB
import streamlit as st
import torch
import torch.nn.functional as F
from torchvision import transforms
from PIL import Image
import numpy as np
import timm
from pathlib import Path
import traceback
# ----------------------
# Page config
# ----------------------
st.set_page_config(
page_title="Pneumonia X-ray Classifier",
layout="centered"
)
st.title("๐Ÿซ Pneumonia Detection from Chest X-ray")
st.write("Upload a chest X-ray image to classify it as **Normal** or **Pneumonia**.")
DEVICE = "cpu"
CLASS_NAMES = ["Normal", "Pneumonia"]
MODEL_PATH = Path(__file__).resolve().parent / "efficientnet_b3_best.pt"
@st.cache_resource
def load_model():
# Recreate model architecture EXACTLY as training
model = timm.create_model(
"efficientnet_b3",
pretrained=False,
num_classes=2 # Normal vs Pneumonia
)
state_dict = torch.load(MODEL_PATH, map_location=DEVICE)
model.load_state_dict(state_dict)
model.to(DEVICE)
model.eval()
return model
model = load_model()
# ----------------------
# File uploader
# ----------------------
uploaded_file = st.file_uploader(
"Upload Chest X-ray Image",
type=["jpg", "jpeg", "png"]
)
if uploaded_file is not None:
try:
# --- Image loading (HF-safe) ---
image = Image.open(uploaded_file)
image = image.convert("RGB")
image = image.resize((300, 300)) # resize
st.image(image, caption="Uploaded X-ray", use_column_width=True)
# --- Preprocess ---
input_tensor = transforms.ToTensor()(image)
input_tensor = transforms.Normalize(
mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225]
)(input_tensor)
input_tensor = input_tensor.unsqueeze(0)
# --- Inference ---
with torch.no_grad():
logits = model(input_tensor)
probs = F.softmax(logits, dim=1).cpu().numpy()[0]
st.subheader("๐Ÿ” Prediction Results")
for i, class_name in enumerate(CLASS_NAMES):
st.write(f"**{class_name}**: {probs[i]*100:.2f}%")
pred_idx = np.argmax(probs)
st.success(f"๐Ÿฉบ Diagnosis: **{CLASS_NAMES[pred_idx]}**")
except Exception as e:
st.error("โŒ Error while processing the image")
st.code(traceback.format_exc())
pred_idx = np.argmax(probs)
st.success(f"๐Ÿฉบ Diagnosis: **{CLASS_NAMES[pred_idx]}**")