Chest X-ray Image Classifier
This repository contains a fine-tuned Vision Transformer (ViT) model for classifying chest X-ray images, utilizing the CheXpert dataset. The model is fine-tuned on the task of classifying various lung diseases from chest radiographs, achieving impressive accuracy in distinguishing between different conditions.
Model Overview
The fine-tuned
Model Code on github is based on the
Vision Transformer (ViT) architecture, which excels in handling image-based tasks by leveraging attention mechanisms for efficient feature extraction. The model was trained on the
CheXpert dataset, which consists of labeled chest X-ray images for detecting diseases such as pneumonia, cardiomegaly, and others.
Performance
- Final Validation Accuracy: 98.46%
- Final Training Loss: 0.1069
- Final Validation Loss: 0.0980
The model achieved a significant accuracy improvement during training, demonstrating its ability to generalize well to unseen chest X-ray images.
Dataset
The dataset used for fine-tuning the model is the CheXpert dataset, which includes chest X-ray images from various patients with multi-label annotations. The data includes frontal and lateral views of the chest for each patient, annotated with labels for various lung diseases.
For more details on the dataset, visit the
CheXpert official website.
Training Details
The model was fine-tuned using the following settings:
- Optimizer: AdamW
- Learning Rate: 3e-5
- Batch Size: 32
- Epochs: 10
- Loss Function: Binary Cross-Entropy with Logits
- Precision: Mixed precision (via
torch.amp)
Usage
Inference
To use the fine-tuned model for inference, simply load the model from Hugging Face's Model Hub and input a chest X-ray image:
1from PIL import Image
2import torch
3from transformers import AutoImageProcessor, AutoModelForImageClassification
4
5# Load model and processor
6processor = AutoImageProcessor.from_pretrained("codewithdark/vit-chest-xray")
7model = AutoModelForImageClassification.from_pretrained("codewithdark/vit-chest-xray")
8
9# Define label columns (class names)
10label_columns = ['Cardiomegaly', 'Edema', 'Consolidation', 'Pneumonia', 'No Finding']
11
12# Step 1: Load and preprocess the image
13image_path = "/content/images.jpeg" # Replace with your image path
14
15# Open the image
16image = Image.open(image_path)
17
18# Ensure the image is in RGB mode (required by most image classification models)
19if image.mode != 'RGB':
20 image = image.convert('RGB')
21 print("Image converted to RGB.")
22
23# Step 2: Preprocess the image using the processor
24inputs = processor(images=image, return_tensors="pt")
25
26# Step 3: Make a prediction (using the model)
27with torch.no_grad(): # Disable gradient computation during inference
28 outputs = model(**inputs)
29
30# Step 4: Extract logits and get the predicted class index
31logits = outputs.logits # Raw logits from the model
32predicted_class_idx = torch.argmax(logits, dim=-1).item() # Get the class index
33
34# Step 5: Map the predicted index to a class label
35# You can also use `model.config.id2label`, but we'll use `label_columns` for this task
36predicted_class_label = label_columns[predicted_class_idx]
37
38# Output the results
39print(f"Predicted Class Index: {predicted_class_idx}")
40print(f"Predicted Class Label: {predicted_class_label}")
41
42'''
43Output :
44Predicted Class Index: 4
45Predicted Class Label: No Finding
46'''
Fine-Tuning
To fine-tune the model on your own dataset, you can follow the instructions in this repo to adapt the code to your dataset and training configuration.
Contributing
We welcome contributions! If you have suggestions, improvements, or bug fixes, feel free to fork the repository and open a pull request.
License
This model is available under the MIT License. See LICENSE for more details.
Acknowledgements
- CheXpert Dataset
- Hugging Face for providing the
transformers library and Model Hub.
Happy coding! 🚀