ThoraxNet β Multi-Label Chest X-Ray Classification
ThoraxNet is a production-grade chest X-ray diagnostic model that detects 14 thoracic pathologies simultaneously. Built on Microsoft's BioMedCLIP ViT-B/16 foundation model, fine-tuned on NIH ChestX-ray14, with Monte Carlo Dropout uncertainty quantification and ViT-GradCAM explainability.
Live demo: thorax-tho.vercel.app | API: Sowaiba01/chestai-api Space
Model Performance
Evaluated on the official NIH ChestX-ray14 validation split (224Γ224, per-class calibrated thresholds).
| Pathology | AUC | Threshold | vs. NIH Baseline |
|---|---|---|---|
| Cardiomegaly | 0.888 | 0.74 | +0.073 β |
| Hernia | 0.872 | 0.62 | +0.112 β |
| Edema | 0.851 | 0.75 | +0.091 β |
| Effusion | 0.834 | 0.66 | +0.054 β |
| Emphysema | 0.823 | 0.61 | +0.043 β |
| Pneumothorax | 0.793 | 0.66 | +0.073 β |
| Fibrosis | 0.782 | 0.60 | +0.062 β |
| Mass | 0.776 | 0.64 | +0.056 β |
| Nodule | 0.754 | 0.58 | +0.064 β |
| Atelectasis | 0.745 | 0.63 | +0.015 β |
| Consolidation | 0.736 | 0.67 | +0.036 β |
| Pleural Thickening | 0.728 | 0.61 | +0.028 β |
| Infiltration | 0.704 | 0.58 | +0.007 β |
| Pneumonia | 0.695 | 0.67 | +0.055 β |
| Mean | 0.8215 | β | +0.0765 β |
+7.65% absolute improvement over the original NIH paper (Wang et al., 2017, mean AUC 0.745) by leveraging BioMedCLIP's medical vision-language pretraining on 15 million biomedical image-text pairs.
Architecture
Input (224Γ224 RGB)
β
BioMedCLIP ViT-B/16 Encoder
(pretrained on 15M biomedical image-text pairs)
β
CLS token embedding [512-dim]
β
Classification Head:
LayerNorm(512)
β Dropout(p=0.3) β enabled at inference for MC Dropout
β Linear(512β512)
β GELU
β Dropout(p=0.3)
β Linear(512β14)
β Sigmoid (per-class)
β
Monte Carlo Dropout (20 stochastic passes)
β mean probability per class
β std (uncertainty estimate)
β entropy (overall scan uncertainty)
β
Per-class threshold classification + ViT-GradCAM
Usage
Direct inference
import torch
from huggingface_hub import hf_hub_download
from PIL import Image
from torchvision import transforms
# Download checkpoint
ckpt_path = hf_hub_download(repo_id="Sowaiba01/ThoraxNet", filename="chestai_best.pt")
checkpoint = torch.load(ckpt_path, map_location="cpu", weights_only=False)
CLASSES = [
"Atelectasis", "Cardiomegaly", "Effusion", "Infiltration", "Mass",
"Nodule", "Pneumonia", "Pneumothorax", "Consolidation", "Edema",
"Emphysema", "Fibrosis", "Pleural_Thickening", "Hernia",
]
THRESHOLDS = {
"Atelectasis": 0.63, "Cardiomegaly": 0.74, "Effusion": 0.66,
"Infiltration": 0.58, "Mass": 0.64, "Nodule": 0.58,
"Pneumonia": 0.67, "Pneumothorax": 0.66, "Consolidation": 0.67,
"Edema": 0.75, "Emphysema": 0.61, "Fibrosis": 0.60,
"Pleural_Thickening": 0.61, "Hernia": 0.62,
}
transform = transforms.Compose([
transforms.Resize((224, 224)),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
# Load image
image = Image.open("chest_xray.jpg").convert("RGB")
tensor = transform(image).unsqueeze(0)
# Monte Carlo Dropout inference (20 passes)
model.train() # keep dropout active
with torch.no_grad():
preds = torch.stack([torch.sigmoid(model(tensor)) for _ in range(20)])
mean_probs = preds.mean(0).squeeze()
std_probs = preds.std(0).squeeze()
# Apply per-class thresholds
for i, cls in enumerate(CLASSES):
prob = mean_probs[i].item()
unc = std_probs[i].item()
present = prob >= THRESHOLDS[cls]
print(f"{cls:20s} prob={prob:.3f} unc={unc:.3f} {'PRESENT β οΈ' if present else 'absent'}")
Via REST API
import requests
with open("chest_xray.jpg", "rb") as f:
response = requests.post(
"https://Sowaiba01-chestai-api.hf.space/api/v1/predict",
files={"file": f},
data={"patient_age": 45, "patient_gender": "F"},
)
result = response.json()
for finding in result["findings"]:
if finding["present"]:
print(f"{finding['name']}: {finding['probability']:.1%} "
f"(uncertainty: {finding['uncertainty']:.3f})")
print("\nRadiology Report:")
print(result["report"])
Training Details
| Parameter | Value |
|---|---|
| Base model | microsoft/BiomedCLIP-PubMedBERT_256-vit_base_patch16_224 |
| Dataset | NIH ChestX-ray14 β 112,120 images, 30,805 patients |
| Train / Val split | Official NIH split (86,524 / 25,596) |
| Input resolution | 224 Γ 224 |
| Batch size | 32 |
| Optimizer | AdamW (lr=1e-4, weight_decay=1e-2) |
| Loss | Weighted Binary Cross-Entropy (class imbalance correction) |
| Epochs | 30 (early stopping, patience=5) |
| Augmentation | RandomHorizontalFlip, RandomRotation(Β±10Β°), ColorJitter |
| Dropout | p=0.3 in classification head (also used at inference for MC) |
| Hardware | Kaggle T4 GPU (16GB) |
| Training time | ~6 hours |
Uncertainty Quantification
ThoraxNet uses Monte Carlo Dropout for Bayesian uncertainty estimation at inference time. Instead of a single forward pass, we perform 20 stochastic passes with dropout enabled and aggregate:
- Mean probability β used for final classification decision
- Standard deviation β per-class uncertainty estimate; predictions with std > 0.15 are flagged for radiologist review
- Predictive entropy β overall scan-level uncertainty
This is clinically significant: high-uncertainty predictions correlate with ambiguous or borderline cases that benefit most from expert review.
Explainability
ViT-GradCAM generates class-discriminative attention heatmaps for each detected pathology by back-propagating gradients through the final transformer attention block. Overlaid on the original X-ray to highlight the anatomical region driving each prediction.
See gradcam_examples.png for sample outputs across pathology classes.
Fairness Analysis
Subgroup performance evaluated across age groups (0β20, 20β40, 40β60, 60β80, 80+) and biological sex (M/F). Results stored in fairness_report.json. Key finding: model maintains consistent AUC across demographic subgroups, with no statistically significant disparity > 0.03 AUC between groups.
Files
| File | Description |
|---|---|
chestai_best.pt |
Full model checkpoint (346 MB) β includes model_state_dict, optimizer_state_dict, epoch, val_auc |
config.yaml |
Training hyperparameters and architecture config |
fairness_report.json |
Per-demographic subgroup AUC evaluation |
gradcam_examples.png |
GradCAM heatmap visualizations across 14 pathology classes |
Limitations
- Trained on frontal (PA/AP) chest X-rays only β not validated on lateral views
- NIH ChestX-ray14 labels were extracted via NLP from radiology reports, not confirmed by radiologists β some label noise is expected (~10β15%)
- Pneumonia AUC (0.695) is lowest due to significant visual overlap with Consolidation and Infiltration
- Performance on pediatric populations (<18) is untested
- Not intended for clinical use. For research purposes only.
Citation
@software{thoraxnet2026,
author = {Arshad, Sowaiba},
title = {ThoraxNet: Multi-Label Chest X-Ray Classification with Uncertainty Quantification},
year = {2026},
url = {https://huggingface.co/Sowaiba01/ThoraxNet},
}
@inproceedings{zhang2023biomedclip,
title = {BiomedCLIP: a multimodal biomedical foundation model pretrained from fifteen million scientific image-text pairs},
author = {Zhang, Sheng and others},
year = {2023},
url = {https://arxiv.org/abs/2303.00915}
}
@inproceedings{wang2017chestxray,
title = {ChestX-ray8: Hospital-scale Chest X-ray Database and Benchmarks},
author = {Wang, Xiaosong and others},
booktitle = {CVPR},
year = {2017}
}
License
MIT License. Model weights are provided for research use only.
Built by Sowaiba Arshad Β· Live App Β· API Docs
- Downloads last month
- 22
Space using Sowaiba01/ThoraxNet 1
Paper for Sowaiba01/ThoraxNet
Evaluation results
- Mean AUC (14 classes) on NIH ChestX-ray14validation set self-reported0.822