ayanahmedkhan's picture
Upload readme.md with huggingface_hub
1e6b95a verified
|
Raw History Blame Contribute Delete
3.4 kB
metadata
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)

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.