Spaces:
Runtime error
Runtime error
| 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" | |
| 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]}**") | |