import os import numpy as np import torch import torch.nn as nn from PIL import Image import timm import albumentations as A from albumentations.pytorch import ToTensorV2 import gradio as gr import matplotlib.pyplot as plt import matplotlib.patches as mpatches # ── config ──────────────────────────────────────────────────────────────────── LABEL_COLS = [ 'No Finding', 'Enlarged Cardiomediastinum', 'Cardiomegaly', 'Lung Opacity', 'Lung Lesion', 'Edema', 'Consolidation', 'Pneumonia', 'Atelectasis', 'Pneumothorax', 'Pleural Effusion', 'Pleural Other', 'Fracture', 'Support Devices' ] COMP_COLS = ['Atelectasis', 'Cardiomegaly', 'Consolidation', 'Edema', 'Pleural Effusion'] # calibrated thresholds from validation set THRESHOLDS = { 'No Finding': 0.80, 'Enlarged Cardiomediastinum': 0.66, 'Cardiomegaly': 0.66, 'Lung Opacity': 0.64, 'Lung Lesion': 0.60, 'Edema': 0.72, 'Consolidation': 0.70, 'Pneumonia': 0.86, 'Atelectasis': 0.72, 'Pneumothorax': 0.72, 'Pleural Effusion': 0.64, 'Pleural Other': 0.90, 'Fracture': 0.50, 'Support Devices': 0.64, } DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu') # ── model ───────────────────────────────────────────────────────────────────── class CheXpertModel(nn.Module): def __init__(self): super().__init__() self.backbone = timm.create_model('densenet121', pretrained=False, num_classes=0, global_pool='avg') feat = self.backbone.num_features self.head = nn.Sequential( nn.BatchNorm1d(feat), nn.Dropout(0.3), nn.Linear(feat, 512), nn.GELU(), nn.BatchNorm1d(512), nn.Dropout(0.15), nn.Linear(512, 14), ) def forward(self, x): return self.head(self.backbone(x)) def load_model(): model = CheXpertModel().to(DEVICE) ckpt = torch.load('model.pth', map_location=DEVICE, weights_only=False) state = ckpt.get('model', ckpt) if any(k.startswith('module.') for k in state.keys()): state = {k.replace('module.', '', 1): v for k, v in state.items() if k != 'n_averaged'} model.load_state_dict(state) model.eval() print(f"model loaded | auc-5: {ckpt.get('auc_5', 'unknown')}") return model transform = A.Compose([ A.Resize(320, 320), A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ToTensorV2(), ]) model = load_model() # ── inference ───────────────────────────────────────────────────────────────── def predict(image): img = np.array(image.convert('RGB')) tensor = transform(image=img)['image'].unsqueeze(0).to(DEVICE) with torch.no_grad(): variants = [tensor, torch.flip(tensor, dims=[-1]), tensor * 0.9, tensor * 1.1] probs = torch.stack([ torch.sigmoid(model(v)) for v in variants ]).mean(0).squeeze().cpu().numpy() results = dict(zip(LABEL_COLS, probs)) positives = [col for col, p in results.items() if p >= THRESHOLDS[col]] # ── build chart ─────────────────────────────────────────────────────────── sorted_items = sorted(results.items(), key=lambda x: -x[1]) labels = [k for k, _ in sorted_items] values = [v for _, v in sorted_items] colors = [] for col, val in sorted_items: thresh = THRESHOLDS[col] if val >= thresh: colors.append('#ef4444') # positive — red elif val >= thresh - 0.08: colors.append('#f97316') # uncertain — orange else: colors.append('#94a3b8') # negative — grey fig, ax = plt.subplots(figsize=(9, 6)) fig.patch.set_facecolor('#0f172a') ax.set_facecolor('#0f172a') bars = ax.barh(labels, values, color=colors, height=0.6, edgecolor='none') # threshold markers for i, (col, val) in enumerate(sorted_items): t = THRESHOLDS[col] ax.plot([t, t], [i - 0.35, i + 0.35], color='white', alpha=0.3, linewidth=1, linestyle='--') ax.set_xlim(0, 1) ax.set_xlabel('probability', color='#94a3b8', fontsize=10) ax.tick_params(colors='#cbd5e1', labelsize=9) ax.spines[:].set_visible(False) ax.xaxis.set_tick_params(color='#334155') for label in ax.get_yticklabels(): col_name = label.get_text() if col_name in COMP_COLS: label.set_color('#60a5fa') else: label.set_color('#cbd5e1') legend_handles = [ mpatches.Patch(color='#ef4444', label='positive'), mpatches.Patch(color='#f97316', label='uncertain'), mpatches.Patch(color='#94a3b8', label='negative'), mpatches.Patch(color='#60a5fa', label='competition label'), ] ax.legend(handles=legend_handles, loc='lower right', facecolor='#1e293b', edgecolor='none', labelcolor='#cbd5e1', fontsize=8) title = 'POSITIVE: ' + ', '.join(positives) if positives else 'No findings flagged' ax.set_title(title, color='#f1f5f9', fontsize=11, pad=12, loc='left') plt.tight_layout() # ── text summary ────────────────────────────────────────────────────────── lines = [] if positives: lines.append('**Flagged findings:**') for col in positives: tag = ' *(competition label)*' if col in COMP_COLS else '' lines.append(f'- {col} — {results[col]:.3f}{tag}') else: lines.append('**No findings flagged above threshold.**') lines.append('\n---') lines.append('*This tool is for research purposes only and is not a medical device.*') summary = '\n'.join(lines) return fig, summary # ── ui ──────────────────────────────────────────────────────────────────────── with gr.Blocks( title='CheXpert Chest X-Ray Classifier', theme=gr.themes.Base( primary_hue='blue', neutral_hue='slate', font=gr.themes.GoogleFont('IBM Plex Mono'), ), css=''' .gradio-container { max-width: 960px; margin: 0 auto; } #title { text-align: center; padding: 24px 0 8px; } #subtitle { text-align: center; color: #94a3b8; margin-bottom: 24px; font-size: 14px; } #disclaimer { font-size: 12px; color: #64748b; text-align: center; margin-top: 8px; } ''' ) as demo: gr.HTML('

CheXpert Chest X-Ray Classifier

') gr.HTML( '

DenseNet-121 · 14 pathologies · AUC-5: 0.897 · ' 'trained on CheXpert

' ) with gr.Row(): with gr.Column(scale=1): image_input = gr.Image(type='pil', label='upload chest x-ray') run_btn = gr.Button('analyse', variant='primary') with gr.Column(scale=2): chart_output = gr.Plot(label='findings') summary_output = gr.Markdown() run_btn.click(fn=predict, inputs=image_input, outputs=[chart_output, summary_output]) gr.HTML( '

⚠️ Research only — not validated for clinical use. ' 'Always consult a qualified radiologist.

' ) if __name__ == '__main__': demo.launch()