Multi-Label Chest X-Ray Classifier

ResNet-50 | PyTorch | NIH ChestX-ray14 | Grad-CAM

An end-to-end medical computer vision project using an ImageNet-pretrained ResNet-50 to classify 14 thoracic conditions from chest X-ray images. The model was trained on the NIH ChestX-ray14 dataset with patient-level data splitting, class-weighted binary cross-entropy, and experiment tracking through Weights & Biases.

Held-Out Test Mean ROC-AUC: 0.823

Live Interactive Demo | GitHub Repository

Model Overview

Property Description
Architecture ResNet-50 with custom classification head
Framework PyTorch
Initialization ImageNet-pretrained weights
Task Multi-label image classification
Dataset NIH ChestX-ray14
Dataset Size 112,120 images from 30,805 patients
Output Classes 14 thoracic conditions
Input Resolution 224 × 224
Loss Function Class-weighted BCEWithLogitsLoss
Optimizer Adam
Training Hardware NVIDIA Tesla T4
Test Mean ROC-AUC 0.823
Model Selection Best validation ROC-AUC, epoch 7

Supported Conditions

The model generates independent prediction scores for 14 thoracic conditions:

  1. Atelectasis
  2. Cardiomegaly
  3. Effusion
  4. Infiltration
  5. Mass
  6. Nodule
  7. Pneumonia
  8. Pneumothorax
  9. Consolidation
  10. Edema
  11. Emphysema
  12. Fibrosis
  13. Pleural Thickening
  14. Hernia

Because this is a multi-label classification task, multiple conditions may be associated with the same X-ray.

The model produces 14 logits, converted into scores using the sigmoid function. These scores are not clinically calibrated probabilities.

Model Architecture

The classifier uses an ImageNet-pretrained ResNet-50 backbone with a custom classification head.

Chest X-Ray (224 × 224 × 3)
             |
     ResNet-50 Backbone
             |
     2048-D Feature Vector
             |
       Linear(2048, 512)
             |
            ReLU
             |
        Dropout(0.3)
             |
        Linear(512, 14)
             |
       14 Output Logits
             |
     Sigmoid (Inference)
             |
     14 Prediction Scores

The model is implemented in PyTorch using the custom ChestXRayNet class.

Dataset and Preprocessing

The model was trained on the NIH ChestX-ray14 dataset, containing 112,120 frontal chest X-ray images from 30,805 patients.

Preprocessing includes:

  • RGB image conversion
  • Resizing to 224 × 224 pixels
  • ImageNet mean and standard deviation normalization
  • Training-time horizontal flipping, brightness/contrast adjustments, and geometric augmentation

Patient-Level Data Splitting

To prevent patient-level data leakage, the dataset was split by patient identifier rather than by individual image.

Split Images Patients
Training 78,566 21,563
Validation 17,063 4,621
Test 16,491 4,621
Total 112,120 30,805

The split uses a 70/15/15 patient-level allocation with random seed 42 and zero patient overlap across partitions.

Images labeled No Finding were retained as negative examples for all 14 conditions.

Training Methodology

Training was performed on Kaggle using an NVIDIA Tesla T4 GPU.

Hyperparameter Value
Epochs 10
Batch Size 32
Learning Rate 0.0001
Optimizer Adam
Loss Weighted BCEWithLogitsLoss
Input Size 224 × 224
Dropout 0.3
Checkpoint Selection Highest validation mean ROC-AUC

Class Imbalance Handling

The NIH ChestX-ray14 dataset contains substantial label imbalance, particularly for rare conditions such as Hernia and Pneumonia.

To address this, positive class weights were calculated from the training set:

positive_weight = negative_samples / positive_samples

These weights were incorporated into BCEWithLogitsLoss to increase the contribution of positive examples from underrepresented classes.

Training and validation metrics were tracked using Weights & Biases.

Evaluation Results

The best checkpoint was selected at epoch 7 using validation mean ROC-AUC.

  • Best Validation Mean ROC-AUC: 0.8315
  • Held-Out Test Mean ROC-AUC: 0.8225
  • Test Images: 16,491
  • Test Patients: 4,621

Per-Class Test ROC-AUC

Condition ROC-AUC
Atelectasis 0.7935
Cardiomegaly 0.8962
Effusion 0.8719
Infiltration 0.6960
Mass 0.8045
Nodule 0.7422
Pneumonia 0.7380
Pneumothorax 0.8775
Consolidation 0.7913
Edema 0.8926
Emphysema 0.9268
Fibrosis 0.7957
Pleural Thickening 0.7947
Hernia 0.8949
Mean 0.8225

ROC-AUC measures the model's ability to distinguish positive and negative examples across classification thresholds. It does not establish clinical diagnostic accuracy, calibration, sensitivity, or specificity at a particular operating threshold.

Model Inference

The trained model checkpoint is available in this repository as best_model.pth.

Inference requires the matching ChestXRayNet architecture and preprocessing pipeline.

For an interactive inference experience, visit the Hugging Face Space.

The deployed application allows users to:

  • Upload a chest X-ray image
  • Generate prediction scores for all 14 conditions
  • View conditions ranked by model score
  • Generate a Grad-CAM heatmap for the highest-scoring condition
  • Visualize the heatmap overlaid on the original image

Model Interpretability

Grad-CAM (Gradient-weighted Class Activation Mapping) is integrated into the model to visualize image regions contributing to a selected class output.

The implementation uses gradients and activations from the final ResNet-50 convolutional block to generate spatial heatmaps.

Grad-CAM visualizations are interpretability aids, not validated medical localization results. Highlighted regions do not necessarily correspond to clinically confirmed pathology.

Intended Use

This model is intended for:

  • Educational and research exploration of medical computer vision
  • Multi-label classification experimentation
  • Transfer learning and class imbalance studies
  • Model interpretability demonstrations
  • Machine learning engineering portfolio evaluation

Limitations and Ethical Considerations

  • The model was evaluated on a held-out partition of NIH ChestX-ray14, not an independent external clinical dataset.
  • NIH ChestX-ray14 labels are derived from radiology reports and may contain annotation noise.
  • Performance may vary across hospitals, imaging equipment, patient populations, and acquisition conditions.
  • Rare conditions have relatively few positive examples, limiting the stability of their performance estimates.
  • Model scores have not been clinically calibrated.
  • No prospective clinical evaluation or regulatory validation has been performed.
  • Grad-CAM heatmaps should not be interpreted as definitive pathology localization.
  • The model should not be used for diagnosis, treatment decisions, or autonomous medical screening.

This model is a research and educational prototype, not a clinically validated medical device.

Project Resources

Author

Samer Ahmed

Developed as a medical computer vision and deep learning engineering project.

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Dataset used to train Samer17/chest-xray-multilabel-classifier

Space using Samer17/chest-xray-multilabel-classifier 1

Collection including Samer17/chest-xray-multilabel-classifier