ayanahmedkhan commited on
Commit
ed17f4f
·
verified ·
1 Parent(s): 1e6b95a

Upload README.md with huggingface_hub

Browse files
Files changed (1) hide show
  1. README.md +427 -1
README.md CHANGED
@@ -1,3 +1,429 @@
1
  ---
2
- license: mit
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3
  ---
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
+ title: "ViT Base Patch16 384 – GI Endoscopy Classifier"
3
+ emoji: "🔬"
4
+ colorFrom: blue
5
+ colorTo: purple
6
+ sdk: pytorch
7
+ sdk_version: "2.0"
8
+ app_file: app.py
9
+ pinned: false
10
+ tags:
11
+ - vision-transformer
12
+ - vit
13
+ - image-classification
14
+ - medical-imaging
15
+ - gastrointestinal
16
+ - endoscopy
17
+ - hyper-kvasir
18
+ - pytorch
19
+ - timm
20
+ - deep-learning
21
+ library_name: timm
22
+ license: other
23
+ language: en
24
+ pipeline_tag: image-classification
25
+ datasets:
26
+ - hyper-kvasir
27
+ metrics:
28
+ - accuracy
29
+ - precision
30
+ - recall
31
+ - f1
32
  ---
33
+
34
+ <div align="center">
35
+
36
+ # 🔬 ViT Base Patch16 384 – GI Endoscopy Classifier
37
+
38
+ **State-of-the-art Vision Transformer for 23-class Gastrointestinal Endoscopy Image Classification**
39
+
40
+ [![PyTorch](https://img.shields.io/badge/PyTorch-2.0+-red?logo=pytorch)](https://pytorch.org)
41
+ [![timm](https://img.shields.io/badge/timm-0.9+-blue)](https://github.com/huggingface/pytorch-image-models)
42
+ [![License](https://img.shields.io/badge/License-Research-green)]()
43
+ [![Accuracy](https://img.shields.io/badge/Test%20Accuracy-93.25%25-brightgreen)]()
44
+
45
+ </div>
46
+
47
+ ---
48
+
49
+ ## 📋 Overview
50
+
51
+ This repository contains a fine-tuned **ViT Base Patch16 384** model for classifying gastrointestinal endoscopy images into 23 anatomical/pathological categories. Trained on the [Hyper-Kvasir](https://datasets.simula.no/hyper-kvasir/) dataset with advanced augmentation techniques including MixUp, Focal Loss, and Test-Time Augmentation (TTA).
52
+
53
+ ### ✨ Key Features
54
+
55
+ | Feature | Description |
56
+ |---------|-------------|
57
+ | 🎯 **High Accuracy** | 93.25% test accuracy with TTA |
58
+ | 🔥 **Modern Architecture** | ViT Base Patch16 @ 384×384 resolution |
59
+ | 📊 **Robust Training** | MixUp, Focal Loss, Label Smoothing, CoarseDropout |
60
+ | ⚡ **Production Ready** | TorchScript traced weights for fast inference |
61
+ | 🧪 **TTA Support** | Test-Time Augmentation for improved predictions |
62
+
63
+ ---
64
+
65
+ ## 📈 Performance Metrics
66
+
67
+ ### Final Results
68
+
69
+ | Metric | Validation (Best) | Test (with TTA) |
70
+ |--------|-------------------|-----------------|
71
+ | **Accuracy** | 92.18% | **93.25%** |
72
+ | **Precision** | – | 92.19% |
73
+ | **Recall** | – | 93.25% |
74
+ | **F1-Score** | – | 92.59% |
75
+
76
+ ### Training Progression
77
+
78
+ | Epoch | Train Acc | Val Acc | Learning Rate | Checkpoint |
79
+ |-------|-----------|---------|---------------|------------|
80
+ | 1 | 50.58% | 81.93% | 4.00e-06 | ✅ |
81
+ | 2 | 67.99% | 86.68% | 6.00e-06 | ✅ |
82
+ | 3 | 74.18% | 87.87% | 8.00e-06 | ✅ |
83
+ | 4 | 74.81% | 88.81% | 1.00e-05 | ✅ |
84
+ | 5 | 77.37% | 89.12% | 1.00e-05 | ✅ |
85
+ | 6 | 77.56% | 89.49% | 9.94e-06 | ✅ |
86
+ | 8 | 80.09% | 90.56% | 9.46e-06 | ✅ |
87
+ | 9 | 80.08% | 90.68% | 9.05e-06 | ✅ |
88
+ | 10 | 80.44% | 90.81% | 8.54e-06 | ✅ |
89
+ | 12 | 82.21% | 91.62% | 7.27e-06 | ✅ |
90
+ | 16 | 85.41% | 91.74% | 4.22e-06 | ✅ |
91
+ | 18 | 84.59% | 92.06% | 2.73e-06 | ✅ |
92
+ | 20 | 86.29% | 92.12% | 1.46e-06 | ✅ |
93
+ | **21** | **85.86%** | **92.18%** | 9.55e-07 | ✅ **Best** |
94
+ | 25 | 86.17% | 92.12% | 0.00e+00 | – |
95
+
96
+ ---
97
+
98
+ ## 🏗️ Model Architecture
99
+
100
+ ```
101
+ ┌─────────────────────────────────────────────────────────────┐
102
+ │ ViT Base Patch16 384 │
103
+ ├─────────────────────────────────────────────────────────────┤
104
+ │ Input: 384 × 384 × 3 (RGB) │
105
+ │ Patch Size: 16 × 16 │
106
+ │ Patches: (384/16)² = 576 patches │
107
+ │ Hidden Dim: 768 │
108
+ │ Layers: 12 Transformer blocks │
109
+ │ Heads: 12 attention heads │
110
+ │ Parameters: 86,108,183 (~86.1M) │
111
+ │ Output: 23 classes (softmax) │
112
+ └─────────────────────────────────────────────────────────────┘
113
+ ```
114
+
115
+ ---
116
+
117
+ ## 🗂️ Dataset: Hyper-Kvasir
118
+
119
+ | Split | Images | Classes |
120
+ |-------|--------|---------|
121
+ | Train | 7,463 | 23 |
122
+ | Validation | 1,599 | 23 |
123
+ | Test | 1,600 | 23 |
124
+ | **Total** | **10,662** | **23** |
125
+
126
+ ### 23 GI Classes
127
+ Anatomical landmarks and pathological findings from upper and lower GI tract endoscopy.
128
+
129
+ ---
130
+
131
+ ## ⚙️ Training Configuration
132
+
133
+ ### Environment
134
+ ```
135
+ PyTorch: 2.x (CUDA 11.8)
136
+ GPU: NVIDIA GPU with ~16GB VRAM
137
+ Python: 3.12
138
+ Platform: Google Colab
139
+ ```
140
+
141
+ ### Dependencies
142
+ ```bash
143
+ pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118
144
+ pip install timm "albumentations>=1.0.0" opencv-python Pillow numpy scikit-learn matplotlib seaborn tqdm
145
+ ```
146
+
147
+ ### Hyperparameters
148
+
149
+ | Parameter | Value |
150
+ |-----------|-------|
151
+ | Model | `vit_base_patch16_384` |
152
+ | Image Size | 384 × 384 |
153
+ | Batch Size | 2 |
154
+ | Effective Batch Size | 16 (8× gradient accumulation) |
155
+ | Epochs | 25 |
156
+ | Base Learning Rate | 1e-5 |
157
+ | Optimizer | AdamW (weight_decay=0.01) |
158
+ | Scheduler | Cosine Annealing + 5-epoch Warmup |
159
+ | Loss | Focal Loss (γ=2.0) + Label Smoothing (0.1) |
160
+ | Mixed Precision | ✅ FP16 (GradScaler) |
161
+ | MixUp | ✅ (α=0.2, p=0.5) |
162
+
163
+ ### Data Augmentation (Albumentations)
164
+
165
+ **Training:**
166
+ ```python
167
+ A.Compose([
168
+ A.Resize(384, 384),
169
+ A.HorizontalFlip(p=0.5),
170
+ A.VerticalFlip(p=0.3),
171
+ A.RandomRotate90(p=0.5),
172
+ A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=15, p=0.5),
173
+ A.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1, p=0.5),
174
+ A.GaussNoise(p=0.3),
175
+ A.CoarseDropout(max_holes=1, max_height=32, max_width=32, p=0.3),
176
+ A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
177
+ ToTensorV2()
178
+ ])
179
+ ```
180
+
181
+ **Validation/Test:**
182
+ ```python
183
+ A.Compose([
184
+ A.Resize(384, 384),
185
+ A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
186
+ ToTensorV2()
187
+ ])
188
+ ```
189
+
190
+ ---
191
+
192
+ ## 🚀 Quick Start
193
+
194
+ ### Installation
195
+
196
+ ```bash
197
+ pip install torch torchvision timm albumentations
198
+ ```
199
+
200
+ ### Inference (TorchScript)
201
+
202
+ ```python
203
+ import torch
204
+ from PIL import Image
205
+ from torchvision import transforms
206
+
207
+ # Load traced model
208
+ model = torch.jit.load("vit_best_traced.pt")
209
+ model.eval()
210
+
211
+ # Preprocessing (must match training)
212
+ preprocess = transforms.Compose([
213
+ transforms.Resize((384, 384)),
214
+ transforms.ToTensor(),
215
+ transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
216
+ ])
217
+
218
+ # Load and classify image
219
+ img = Image.open("endoscopy_image.jpg").convert("RGB")
220
+ tensor = preprocess(img).unsqueeze(0)
221
+
222
+ with torch.no_grad():
223
+ logits = model(tensor)
224
+ probs = logits.softmax(dim=1)
225
+ confidence, pred_class = probs.max(dim=1)
226
+
227
+ print(f"Predicted class: {pred_class.item()}")
228
+ print(f"Confidence: {confidence.item():.2%}")
229
+ ```
230
+
231
+ ### Inference with Test-Time Augmentation (TTA)
232
+
233
+ ```python
234
+ import torch
235
+ from PIL import Image
236
+ from torchvision import transforms
237
+
238
+ model = torch.jit.load("vit_best_traced.pt")
239
+ model.eval()
240
+
241
+ preprocess = transforms.Compose([
242
+ transforms.Resize((384, 384)),
243
+ transforms.ToTensor(),
244
+ transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
245
+ ])
246
+
247
+ def predict_with_tta(model, tensor):
248
+ """Test-Time Augmentation: average predictions across flips"""
249
+ with torch.no_grad():
250
+ # Original
251
+ pred1 = model(tensor).softmax(dim=1)
252
+ # Horizontal flip
253
+ pred2 = model(torch.flip(tensor, [3])).softmax(dim=1)
254
+ # Vertical flip
255
+ pred3 = model(torch.flip(tensor, [2])).softmax(dim=1)
256
+ # Average
257
+ return (pred1 + pred2 + pred3) / 3.0
258
+
259
+ img = Image.open("endoscopy_image.jpg").convert("RGB")
260
+ tensor = preprocess(img).unsqueeze(0)
261
+
262
+ probs = predict_with_tta(model, tensor)
263
+ confidence, pred_class = probs.max(dim=1)
264
+
265
+ print(f"Predicted class (TTA): {pred_class.item()}")
266
+ print(f"Confidence: {confidence.item():.2%}")
267
+ ```
268
+
269
+ ### Batch Inference
270
+
271
+ ```python
272
+ import torch
273
+ from PIL import Image
274
+ from torchvision import transforms
275
+ from pathlib import Path
276
+
277
+ model = torch.jit.load("vit_best_traced.pt")
278
+ model.eval()
279
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
280
+ model = model.to(device)
281
+
282
+ preprocess = transforms.Compose([
283
+ transforms.Resize((384, 384)),
284
+ transforms.ToTensor(),
285
+ transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
286
+ ])
287
+
288
+ def classify_batch(image_paths, batch_size=8):
289
+ results = []
290
+ for i in range(0, len(image_paths), batch_size):
291
+ batch_paths = image_paths[i:i+batch_size]
292
+ tensors = []
293
+ for path in batch_paths:
294
+ img = Image.open(path).convert("RGB")
295
+ tensors.append(preprocess(img))
296
+
297
+ batch = torch.stack(tensors).to(device)
298
+ with torch.no_grad():
299
+ probs = model(batch).softmax(dim=1)
300
+ confidences, preds = probs.max(dim=1)
301
+
302
+ for path, pred, conf in zip(batch_paths, preds, confidences):
303
+ results.append({
304
+ "file": str(path),
305
+ "class": pred.item(),
306
+ "confidence": conf.item()
307
+ })
308
+ return results
309
+
310
+ # Example usage
311
+ image_folder = Path("./test_images")
312
+ image_paths = list(image_folder.glob("*.jpg"))
313
+ results = classify_batch(image_paths)
314
+ ```
315
+
316
+ ---
317
+
318
+ ## 📁 Repository Structure
319
+
320
+ ```
321
+ .
322
+ ├── vit_best_traced.pt # TorchScript traced weights (best checkpoint)
323
+ ├── README.md # This file
324
+ └��─ class_mapping.json # (Optional) Class index to name mapping
325
+ ```
326
+
327
+ ---
328
+
329
+ ## 🔧 Advanced: Custom Training
330
+
331
+ ### Focal Loss Implementation
332
+
333
+ ```python
334
+ class FocalLoss(nn.Module):
335
+ def __init__(self, alpha=1, gamma=2, reduction='mean'):
336
+ super().__init__()
337
+ self.alpha = alpha
338
+ self.gamma = gamma
339
+ self.reduction = reduction
340
+
341
+ def forward(self, inputs, targets):
342
+ ce_loss = F.cross_entropy(inputs, targets, reduction='none')
343
+ pt = torch.exp(-ce_loss)
344
+ focal_loss = self.alpha * (1 - pt) ** self.gamma * ce_loss
345
+
346
+ if self.reduction == 'mean':
347
+ return focal_loss.mean()
348
+ return focal_loss.sum() if self.reduction == 'sum' else focal_loss
349
+ ```
350
+
351
+ ### MixUp Implementation
352
+
353
+ ```python
354
+ def mixup_data(x, y, alpha=0.2):
355
+ lam = np.random.beta(alpha, alpha) if alpha > 0 else 1
356
+ batch_size = x.size(0)
357
+ index = torch.randperm(batch_size).to(x.device)
358
+
359
+ mixed_x = lam * x + (1 - lam) * x[index]
360
+ y_a, y_b = y, y[index]
361
+ return mixed_x, y_a, y_b, lam
362
+
363
+ def mixup_criterion(criterion, pred, y_a, y_b, lam):
364
+ return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)
365
+ ```
366
+
367
+ ---
368
+
369
+ ## ⚠️ Limitations & Responsible Use
370
+
371
+ > **⚕️ Medical Disclaimer**
372
+ >
373
+ > This model is a **research artifact** and is **NOT** a regulated medical device. It should **NOT** be used for clinical diagnosis without proper validation and regulatory approval.
374
+
375
+ ### Known Limitations
376
+ - Trained on Hyper-Kvasir dataset; may not generalize to other endoscopy equipment or populations
377
+ - Best performance requires 384×384 input resolution
378
+ - TTA improves accuracy but increases inference time 3×
379
+
380
+ ### Recommended Use
381
+ - ✅ Research and educational purposes
382
+ - ✅ Preliminary screening with human oversight
383
+ - ✅ Benchmark for GI image classification
384
+ - ❌ Standalone clinical diagnosis
385
+ - ❌ Life-critical medical decisions
386
+
387
+ ---
388
+
389
+ ## 📚 Citation
390
+
391
+ If you use this model in your research, please cite:
392
+
393
+ ```bibtex
394
+ @misc{vit_gi_endoscopy_2025,
395
+ author = {Ayan Ahmed Khan},
396
+ title = {ViT Base Patch16 384 for GI Endoscopy Classification},
397
+ year = {2025},
398
+ publisher = {Hugging Face},
399
+ url = {https://huggingface.co/ayanahmedkhan/VIT-gi-endoscopy-classifier}
400
+ }
401
+ ```
402
+
403
+ ### Related Work
404
+ - [An Image is Worth 16x16 Words (ViT)](https://arxiv.org/abs/2010.11929)
405
+ - [Hyper-Kvasir Dataset](https://datasets.simula.no/hyper-kvasir/)
406
+ - [timm Library](https://github.com/huggingface/pytorch-image-models)
407
+
408
+ ---
409
+
410
+ ## 📝 Changelog
411
+
412
+ | Date | Version | Changes |
413
+ |------|---------|---------|
414
+ | 2025-12-29 | 1.0.0 | Initial release with traced weights and full documentation |
415
+
416
+ ---
417
+
418
+ ## 📬 Contact
419
+
420
+ - **Author:** Ayan Ahmed Khan
421
+ - **Hugging Face:** [ayanahmedkhan](https://huggingface.co/ayanahmedkhan)
422
+
423
+ ---
424
+
425
+ <div align="center">
426
+
427
+ **Made with ❤️ for Medical AI Research**
428
+
429
+ </div>