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:
- Atelectasis
- Cardiomegaly
- Effusion
- Infiltration
- Mass
- Nodule
- Pneumonia
- Pneumothorax
- Consolidation
- Edema
- Emphysema
- Fibrosis
- Pleural Thickening
- 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.