Projects06 of 15
Glass-Box ViT
A Vision Transformer rebuilt from scratch and verified layer by layer against a pretrained ViT-Base
The question
Vision Transformers are usually used through a library, so their insides stay a black box. The aim was to rebuild one from basic layers and prove it is the same function as a production pretrained model: not merely that it picks the same top-1 class, but that it is equal at every layer. Then to use it to measure which design choices matter.
What was built
- An inspectable ViT: patch projection, learned positions, a CLS token, packed QKV with a hand-written softmax(QKᵀ/√d)V, and pre-norm blocks. An optional analysis mode returns the attention maps and hidden states.
- A strict weight mapping from the pinned timm checkpoint
vit_base_patch16_224.augreg_in21k_ft_in1k, with its revision and SHA-256 verified. It refuses any missing key or shape mismatch. - Tooling: a layer-by-layer comparison tool, a resumable training pipeline with checkpoint-reload proofs, benchmarks, and ablation runners with frozen protocols.
- A Colab T4 notebook for the GPU studies.
It matches the pretrained model exactly
| Check | Result |
|---|---|
| Pretrained tensors mapped | 152 of 152 (86,567,656 parameters) |
| Layer-by-layer equivalence | 100 of 100 stages, max difference 0.0, on an Apple M1 CPU and a Colab Tesla T4 |
| Inference, 224 px, batch 1 | 81.0 ms on the M1 CPU, 15.0 ms on the T4, against 14.6 ms for timm’s fused attention |
What matters in a ViT
CIFAR-10 ablations over three paired seeds that share identical splits, with each protocol frozen and committed before its run:
| Choice | Validation accuracy |
|---|---|
| Patch size 4, 8, 16 at width 96 | 71.6%, 63.4%, 55.9% |
| Patch size 4, 8, 16 at width 192 | 75.3%, 68.7%, 59.3% |
| Learned positions against none | 63.4% against 56.0%; learned won in every seed |
| Mean-patch against CLS pooling | 65.2% against 63.4%; mean won in every seed |
| 4, 8, 12, 16 heads at fixed width | 63.8%, 63.4%, 63.0%, 62.9%; no clear difference |
These are validation accuracies of small models trained from scratch for 20 epochs. The test split was never used.
Transfer to retinal images
A RetinaMNIST transfer study, as a research and education exercise rather than a clinical tool. Quadratic weighted kappa on the official 400-image test set:
- linear probe: 0.750;
- partial fine-tuning: 0.775 ± 0.038;
- full fine-tuning: 0.773 ± 0.038.
Neither fine-tuning method beat the linear probe reliably. Models were selected on validation data only, and test metrics were recomputed from saved predictions with bootstrap 95% intervals.
How the numbers were kept honest
- Every trained checkpoint was reloaded and shown to reproduce its saved validation predictions.
- Two GPU-only bugs, a device mismatch in that reload check and a CPU data-loading bottleneck, were found and fixed before the GPU numbers were trusted.
- Last-layer attention maps are shown as diagnostics, not explanations. CLS attention concentrates on a few background patches, a known ViT artifact.
- A clean clone installs and passes all 34 tests.
What I learned
Matching a pretrained model exactly comes down to small details: the LayerNorm epsilon, the pre-norm order, the QKV split order, and turning off fused attention when comparing kernels. In this small-data regime, patch size and position information mattered far more than the number of heads.
Built with
Python 3.12, PyTorch, timm, torchvision, scikit-learn, MedMNIST, matplotlib, Google Colab (T4) and pytest. A solo project: design, implementation, experiments and write-up. The code is MIT-licensed; the checkpoint and datasets keep their own terms. The demo’s photos are scikit-image sample images, in the public domain.
Read the full report, the experiment log, or open the Colab notebook.