ayanahmedkhan's picture
Upload readme.md with huggingface_hub
1e6b95a verified
|
Raw History Blame Contribute Delete
3.4 kB
---
title: "ViT Base Patch16 384 – GI Endoscopy Classifier"
tags:
- vision-transformer
- vit
- image-classification
- medical-imaging
- gastrointestinal
- pytorch
- timm
library_name: timm
license: other
language: en
pipeline_tag: image-classification
---
## ViT Base (384px) for GI Endoscopy
23-class gastrointestinal endoscopy classifier built on ViT Base Patch16 384. Trained with timm, strong augmentation (MixUp + Albumentations), and evaluated with TTA. Includes traced TorchScript weights `vit_best_traced.pt` for drop-in PyTorch inference.
### Highlights
- Val: **92.18%** (epoch 21 best); Test (TTA): **93.25% Acc / 92.19% Prec / 93.25% Rec / 92.59% F1**
- Robust pipeline: MixUp, CoarseDropout, and TTA for generalization
- Ready to serve: TorchScript traced weights; standard 384x384 ImageNet normalization
- Scope: 23 GI classes; input RGB 384x384
### Performance
| Split | Accuracy | Precision | Recall | F1 | Notes |
| --- | --- | --- | --- | --- | --- |
| Validation (best) | 92.18% | – | – | – | Best at epoch 21 |
| Test (TTA) | 93.25% | 92.19% | 93.25% | 92.59% | 800 batches, TTA enabled |
### Model Card
- Backbone: `vit_base_patch16_384` (≈86.1M params)
- Classes: 23
- Input: 384x384 RGB, mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]
- Training recipe:
- Epochs: 25; batch size: 2
- LR: 1e-5 base with warmup schedule
- Augment: MixUp; Albumentations with CoarseDropout
- Eval: TTA on validation/test; best checkpoint saved each epoch (best at epoch 21)
- Data scale (from logs): 7,463 train, 1,599 val, 1,600 test images across 23 classes
### Inference (PyTorch TorchScript)
```python
import torch
from PIL import Image
from torchvision import transforms
model = torch.jit.load("vit_best_traced.pt")
model.eval()
preprocess = transforms.Compose([
transforms.Resize((384, 384)),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
img = Image.open("path/to/image.jpg").convert("RGB")
tensor = preprocess(img).unsqueeze(0)
with torch.no_grad():
probs = model(tensor).softmax(dim=1)
pred = probs.argmax(dim=1).item()
print("Predicted class index:", pred)
```
### Training Log (key checkpoints)
- Epoch 8: Val Acc 90.56%
- Epoch 12: Val Acc 91.62%
- Epoch 16: Val Acc 91.74%
- Epoch 18: Val Acc 92.06%
- Epoch 20: Val Acc 92.12%
- Epoch 21: **Val Acc 92.18%** (best)
- Test (TTA): Acc 93.25%, Precision 92.19%, Recall 93.25%, F1 92.59%
### Files
- `vit_best_traced.pt` — TorchScript traced weights for ViT Base Patch16 384 (best checkpoint)
### Notes & Responsible Use
- Trainer: custom `AdvancedMemoryEfficientTrainer` with MixUp and Albumentations TTA.
- Warnings during training: Albumentations recommends `Affine` over `ShiftScaleRotate`; CoarseDropout args adjusted. Metrics unaffected.
- Medical context: research-only; not a regulated medical device. Validate clinically and keep a human in the loop.
### Reproduce / Fine-tune Tips
- Keep preprocessing (384 resize + ImageNet normalization) consistent to match reported metrics.
- Start from `vit_base_patch16_384` in timm; use MixUp and similar augmentations.
- Enable TTA for best accuracy; disable for faster inference.
### Changelog
- 2025-12-29: Initial public release with traced weights and metrics.