rsna / src /streamlit_app.py
Dongqing Geng
Update src/streamlit_app.py
d842a9b verified
Raw
History Blame Contribute Delete
7.34 kB
import streamlit as st
import torch
import numpy as np
import cv2
from PIL import Image
from transformers import AutoModelForImageClassification, AutoImageProcessor
from torchvision import transforms
from pytorch_grad_cam import GradCAM
from pytorch_grad_cam.utils.image import show_cam_on_image
from pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget
import torch.nn.functional as F
# ==========================================
# 1. 页面配置与标题
# ==========================================
st.set_page_config(
page_title="RSNA Intracranial Hemorrhage AI",
page_icon="🧠",
layout="wide"
)
st.title("🧠 AI-Assisted Intracranial Hemorrhage Detection")
st.markdown("""
**Workflow:** `Input CT Scan` $\\rightarrow$ `ResNet50 Inference` $\\rightarrow$ `Grad-CAM XAI` $\\rightarrow$ `Radiologist Review`
""")
# ==========================================
# 2. 模型加载 (使用 @st.cache_resource 加速)
# ==========================================
MODEL_ID = "dongqinggeng/rsna" # 你的 Hugging Face Model ID
@st.cache_resource
def load_model_and_processor():
try:
# 加载模型
model = AutoModelForImageClassification.from_pretrained(MODEL_ID)
model.eval()
# 尝试加载配套的 processor,如果失败则使用默认逻辑
try:
processor = AutoImageProcessor.from_pretrained(MODEL_ID)
except:
processor = None
return model, processor
except Exception as e:
st.error(f"Error loading model from Hugging Face: {e}")
return None, None
model, processor = load_model_and_processor()
# 定义标签映射 (确保顺序与训练时一致)
LABELS = ['epidural', 'intraparenchymal', 'intraventricular', 'subarachnoid', 'subdural', 'any']
id2label = {i: label for i, label in enumerate(LABELS)}
# ==========================================
# 3. 辅助类与函数
# ==========================================
# Grad-CAM 需要的 Wrapper (适配 HF 模型输出)
class HuggingFaceModelWrapper(torch.nn.Module):
def __init__(self, model):
super(HuggingFaceModelWrapper, self).__init__()
self.model = model
def forward(self, x):
return self.model(x).logits
# 图像预处理 (必须与训练时保持一致!)
def process_image(image):
# 定义变换
transform = transforms.Compose([
transforms.Resize((224, 224)), # 如果你之前是 128,这里用 128,ResNet 默认通常适配 224 效果更好
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])
# 确保是 RGB
image = image.convert("RGB")
input_tensor = transform(image).unsqueeze(0) # Add batch dimension
return input_tensor, image
# 生成 Grad-CAM
def generate_gradcam(model, input_tensor, target_layer):
cam = GradCAM(model=HuggingFaceModelWrapper(model), target_layers=[target_layer])
# target=None 表示解释概率最高的那个类
grayscale_cam = cam(input_tensor=input_tensor, targets=None)
return grayscale_cam[0, :]
# ==========================================
# 4. 主界面逻辑
# ==========================================
# --- 侧边栏:上传图片 ---
st.sidebar.header("1. Input Data")
uploaded_file = st.sidebar.file_uploader("Upload a CT Slice (PNG/JPG/DICOM)", type=["png", "jpg", "jpeg"])
# 提供示例图片下载链接或按钮 (可选)
if st.sidebar.button("Load Demo Image (Simulated)"):
# 这里只是为了演示,实际应用中你可以放一个示例文件的 URL
st.sidebar.info("Please upload a real file to test.")
# --- 主区域 ---
col1, col2 = st.columns([1, 1.5])
if uploaded_file is not None and model is not None:
# A. 显示原图
image_pil = Image.open(uploaded_file)
with col1:
st.subheader("Original CT Scan")
st.image(image_pil, use_container_width=True, caption="Input Slice")
# B. 推理 (Inference)
with st.spinner("Running AI Model & Generating Explanations..."):
# 预处理
input_tensor, image_rgb_pil = process_image(image_pil)
# 预测
with torch.no_grad():
outputs = model(input_tensor)
probs = torch.sigmoid(outputs.logits).cpu().numpy()[0]
# 找到概率最高的特定类型 (排除 'any')
# labels 顺序: epidural, intraparenchymal, intraventricular, subarachnoid, subdural, any
# any 是 index 5
specific_probs = probs[:5]
top_idx = np.argmax(specific_probs)
top_label = LABELS[top_idx]
top_prob = specific_probs[top_idx]
any_prob = probs[5]
# C. XAI: Grad-CAM
# 获取 ResNet 最后一层
target_layer = model.resnet.encoder.stages[-1].layers[-1]
cam_mask = generate_gradcam(model, input_tensor, target_layer)
# 图像叠加处理
# 1. 归一化原图用于显示
img_np = np.array(image_rgb_pil)
img_np = img_np.astype(np.float32) / 255.0
# 2. 叠加
visualization = show_cam_on_image(img_np, cam_mask, use_rgb=True)
# D. 显示结果
with col2:
st.subheader("XAI Output (Grad-CAM)")
st.image(visualization, use_container_width=True, caption=f"Model Focus Area (Red = High Attention)")
# E. 详细报告与指标
st.divider()
st.header("📝 AI Analysis Report")
# 动态生成的“放射科医生评论”
if any_prob > 0.5:
status_color = "red"
status_text = "Hemorrhage Detected"
confidence_text = f"High Confidence ({any_prob:.2%})"
review_text = (
f"**AI Findings:** The model detected specific features consistent with **{top_label} hemorrhage** "
f"(Probability: {top_prob:.2%}).\n\n"
f"**XAI Localization:** The Grad-CAM heatmap highlights a region of interest. "
f"Please verify if this corresponds to a hyperdense area in the brain parenchyma or extra-axial space."
)
else:
status_color = "green"
status_text = "No Hemorrhage Detected"
confidence_text = f"({1-any_prob:.2%} sure)"
review_text = "**AI Findings:** No significant signs of intracranial hemorrhage were detected. Heatmap shows diffuse or non-specific activation."
# 使用 Metrics 展示
m1, m2, m3 = st.columns(3)
m1.metric("Overall Prediction", status_text, delta=confidence_text, delta_color="inverse" if any_prob > 0.5 else "normal")
m2.metric("Primary Subtype", top_label.capitalize() if any_prob > 0.3 else "N/A", f"{top_prob:.2%}")
st.markdown(f"""
> **Radiologist Review Note:** > {review_text}
""")
# F. 详细概率条形图
st.subheader("Detailed Class Probabilities")
st.bar_chart({label: prob for label, prob in zip(LABELS, probs)})
else:
# 初始欢迎界面
st.info("👈 Please upload a CT image from the sidebar to start the analysis.")
st.markdown("### How to interpret the heatmap?")
st.markdown("""
* **Red Areas**: Regions that contributed most to the AI's decision (High Importance).
* **Blue Areas**: Regions the AI ignored.
* *Note: If the heatmap highlights the skull or background, the prediction might be an artifact.*
""")