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.* """)