Team Ai
Apppublic

sandy45/ChestViT-Explainable-XRay-AI

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
App README

Multi-Task Vision Transformer for Chest X-Ray Disease Classification & Explainability

<div align="center">

Python PyTorch HuggingFace Gradio MLflow

ViT-Base-16 ยท 14-Disease Multi-Label Classification ยท Attention Rollout XAI

Fine-tuned on NIH ChestX-ray14 (112,120 frontal chest X-rays)

</div>


๐ŸŽฏ Overview

This project implements an explainable medical AI system for automated chest X-ray analysis. Unlike standard classification models, this system simultaneously:

  1. 1.Predicts 14 diseases in parallel (multi-label, not multi-class)
  2. 2.Shows WHERE in the X-ray the model is looking via Attention Rollout
  3. 3.Compares against published NIH baselines (AUC-ROC per class)
  4. 4.Serves a live demo via a Gradio dashboard

The core insight: Vision Transformers (ViTs) divide images into 16ร—16 patches and route information through 12 attention layers. Attention Rollout traces this information flow back to the input patches โ€” telling us exactly which lung regions drove each disease prediction.


๐Ÿ— Architecture

Chest X-Ray Input (PNG, 1024ร—1024)
           โ”‚
           โ–ผ
   โ”Œโ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”
   โ”‚  CLAHE Preprocessing  โ”‚   โ† Contrast Limited Adaptive Histogram Equalization
   โ”‚  + Albumentations     โ”‚   โ† Radiologically-realistic augmentation
   โ””โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”˜
           โ”‚ 224ร—224ร—3
           โ–ผ
   โ”Œโ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”
   โ”‚          ViT-Base-16                      โ”‚
   โ”‚  (google/vit-base-patch16-224-in21k)      โ”‚
   โ”‚                                           โ”‚
   โ”‚  โ”Œโ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”  โ”‚
   โ”‚  โ”‚  196 Patches (14ร—14 grid, 16px/ea) โ”‚  โ”‚
   โ”‚  โ”‚  + [CLS] token = 197 total tokens  โ”‚  โ”‚
   โ”‚  โ””โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”˜  โ”‚
   โ”‚              โ†“                            โ”‚
   โ”‚  12 Transformer Layers                    โ”‚
   โ”‚  (12 heads ร— 64 dim = 768 hidden dim)    โ”‚
   โ”‚              โ†“                            โ”‚
   โ”‚  [CLS] token โ†’ Dropout โ†’ Linear(768โ†’14)  โ”‚
   โ””โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”˜
           โ”‚
           โ”œโ”€โ”€โ”€ Logits โ†’ Sigmoid โ†’ 14 disease probabilities
           โ”‚
           โ””โ”€โ”€โ”€ Attention weights (12 layers ร— 12 heads)
                         โ”‚
                         โ–ผ
              Attention Rollout Algorithm
              (14ร—14 patch attention map)
                         โ”‚
                         โ–ผ
              224ร—224 heatmap overlay on X-ray

๐Ÿ“Š Results

DiseaseViT AUCNIH Baselineฮ” AUC
Atelectasisโ€”0.7003โ€”
Cardiomegalyโ€”0.8100โ€”
Effusionโ€”0.7585โ€”
Infiltrationโ€”0.6614โ€”
Massโ€”0.6933โ€”
Noduleโ€”0.6689โ€”
Pneumoniaโ€”0.6580โ€”
Pneumothoraxโ€”0.7993โ€”
Consolidationโ€”0.7032โ€”
Edemaโ€”0.8052โ€”
Emphysemaโ€”0.8330โ€”
Fibrosisโ€”0.7859โ€”
Pleural_Thickeningโ€”0.6835โ€”
Herniaโ€”0.8717โ€”
MACRO AVERAGEโ€”0.7523โ€”

Results will populate after training. NIH baseline from Wang et al. (2017).


๐Ÿš€ Quick Start

1. Install Dependencies

bash
# Create virtual environment
python -m venv venv
venv\Scripts\activate   # Windows
# source venv/bin/activate  # Linux/macOS

# Install dependencies
pip install -r requirements.txt

2. Download Dataset

bash
# First: set up Kaggle API credentials
# 1. Go to https://www.kaggle.com โ†’ Account โ†’ Settings โ†’ Create New API Token
# 2. Place kaggle.json at: C:\Users\<YourName>\.kaggle\kaggle.json

# Then download NIH ChestX-ray14 (~42 GB)
python data/download.py

3. Run Unit Tests

bash
python -m pytest tests/ -v

4. Train the Model

bash
# Full training (5 epochs, ~8-12 hours on RTX 3050)
python training/train.py

# Quick smoke test (20% of data)
# Edit config/config.yaml โ†’ dataset.train_fraction: 0.2
python training/train.py

Monitor training in real-time:

bash
mlflow ui --backend-store-uri ./experiments/mlflow
# Open http://localhost:5000

5. Launch Demo

bash
# With trained model
python app/gradio_app.py

# DEMO MODE (random weights, for UI preview only)
set DEMO_MODE=1   # Windows
python app/gradio_app.py

๐Ÿ“ Project Structure

.
โ”œโ”€โ”€ config/
โ”‚   โ””โ”€โ”€ config.yaml              # All hyperparameters and paths
โ”œโ”€โ”€ data/
โ”‚   โ”œโ”€โ”€ download.py              # Kaggle API dataset download
โ”‚   โ”œโ”€โ”€ preprocessing.py         # CLAHE + Albumentations pipeline
โ”‚   โ”œโ”€โ”€ dataset.py               # ChestXrayDataset + DataLoaders
โ”‚   โ””โ”€โ”€ raw/                     # Downloaded dataset (not in git)
โ”œโ”€โ”€ models/
โ”‚   โ””โ”€โ”€ vit_model.py             # ViT-Base-16 with multi-label head
โ”œโ”€โ”€ explainability/
โ”‚   โ””โ”€โ”€ attention_rollout.py     # Attention Rollout algorithm
โ”œโ”€โ”€ training/
โ”‚   โ”œโ”€โ”€ losses.py                # Weighted BCE + Focal Loss
โ”‚   โ”œโ”€โ”€ train.py                 # Training loop (mixed precision, MLflow)
โ”‚   โ””โ”€โ”€ evaluate.py              # AUC-ROC per class, ROC plots
โ”œโ”€โ”€ app/
โ”‚   โ””โ”€โ”€ gradio_app.py            # Gradio dashboard
โ”œโ”€โ”€ tests/
โ”‚   โ””โ”€โ”€ test_modules.py          # Unit tests (no dataset required)
โ”œโ”€โ”€ checkpoints/                 # Saved model weights (not in git)
โ”œโ”€โ”€ results/                     # ROC curves, metrics CSV
โ”œโ”€โ”€ experiments/
โ”‚   โ””โ”€โ”€ mlflow/                  # MLflow tracking database
โ”œโ”€โ”€ config_loader.py             # YAML config loader
โ””โ”€โ”€ requirements.txt

โš™๏ธ Configuration

All settings are in `config/config.yaml`. Key RTX 3050 settings:

yaml
training:
  batch_size: 8                    # Fits in 4 GB VRAM
  gradient_accumulation_steps: 4   # Effective batch = 32
  mixed_precision: true            # fp16 โ€” mandatory for 4 GB VRAM
  num_epochs: 5

model:
  name: "google/vit-base-patch16-224-in21k"
  gradient_checkpointing: true     # Saves ~30% VRAM

dataset:
  train_fraction: 1.0              # Set 0.2 for quick smoke test

๐Ÿ”ฅ Attention Rollout: Why Not Grad-CAM?

MethodGrad-CAMAttention Rollout
Designed forCNNsTransformers
Spatial resolutionDepends on last conv layer14ร—14 patch grid
Accounts for skip connectionsNoYes (identity matrix)
Computational costRequires backward passForward pass only
ViT-specificNoYes

Attention Rollout (Abnar & Zuidema, 2020) is mathematically derived from the transformer's own attention mechanism, making it the correct tool for ViT explainability.


๐Ÿฅ Clinical Context

โš ๏ธ This is a research/educational project, NOT a medical device. Results should not be used for clinical diagnosis without radiologist review.

The NIH ChestX-ray14 dataset has known limitations (Rajpurkar et al., 2018, and others). AUC-ROC is the clinically relevant metric because:

  • โ€”Accuracy is misleading with class imbalance (>53% "No Finding")
  • โ€”AUC measures discriminative ability across all thresholds
  • โ€”Radiologists can set their own confidence threshold per clinical context

๐Ÿ“š References

  1. 1.Wang, X. et al. (2017). ChestX-ray8: Hospital-scale Chest X-ray Database and Benchmarks. CVPR.
  2. 2.Dosovitskiy, A. et al. (2021). An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale. ICLR.
  3. 3.Abnar, S. & Zuidema, W. (2020). Quantifying Attention Flow in Transformers. arXiv:2005.00928.
  4. 4.Rajpurkar, P. et al. (2017). CheXNet: Radiologist-Level Pneumonia Detection on Chest X-Rays with Deep Learning. arXiv:1711.05225.

๐Ÿ›  Tech Stack

ToolVersionPurpose
PyTorch2.1+Training framework
HuggingFace Transformers4.37+ViT-Base-16 backbone
OpenCV4.9+CLAHE preprocessing
Albumentations1.3+Image augmentation
scikit-learn1.4+AUC-ROC metrics
MLflow2.10+Experiment tracking
Gradio4.xDemo dashboard
Kaggle API1.6+Dataset download