A learner-friendly implementation of a Vision Transformer (ViT) built from scratch using PyTorch and trained on the MNIST handwritten digit dataset.
The goal of this project is not to use a pre-built Vision Transformer implementation, but to understand the architecture by implementing its major components manually with PyTorch building blocks.
By going through vision_transformer.ipynb, you will understand how an image is transformed into a sequence of tokens and processed by a Transformer.
The notebook covers:
- Image preprocessing with
torchvision.transforms - Loading MNIST with
torchvision.datasets - Creating PyTorch
DataLoaders - Splitting an image into non-overlapping patches
- Patch projection using
Conv2d - Learnable
[CLS]token - Learnable positional embeddings
- Multi-Head Self-Attention
- Layer Normalization
- Residual/skip connections
- Transformer feed-forward MLP
- Stacking multiple Transformer encoder blocks
- Classification using the
[CLS]token - Training with Adam and Cross-Entropy Loss
- Model evaluation
- Visualizing predictions
The complete architecture implemented in the notebook is:
MNIST IMAGE
28 × 28 × 1
│
▼
┌─────────────────────┐
│ Patch Embedding │
│ Conv2d: 7 × 7 │
│ stride = 7 │
└─────────────────────┘
│
▼
16 Patch Tokens
16 × 32
│
│
+ Learnable [CLS]
│
▼
17 × 32 Tokens
│
│
+ Positional Embedding
│
▼
17 × 32 Tokens
│
▼
┌────────────────────────────┐
│ Transformer Encoder × 4 │
│ │
│ LayerNorm │
│ ↓ │
│ Multi-Head Attention │
│ ↓ │
│ Residual Connection │
│ ↓ │
│ LayerNorm │
│ ↓ │
│ MLP: 32 → 64 → 32 │
│ ↓ │
│ Residual Connection │
└────────────────────────────┘
│
▼
[CLS] Representation
1 × 32
│
▼
┌────────────────┐
│ Classification │
│ 32 → 10 │
└────────────────┘
│
▼
MNIST Digit Prediction
0 ... 9
One of the most important parts of implementing a Vision Transformer is understanding how the tensor shape changes.
For a batch size of 64:
Input image
[64, 1, 28, 28]
│
▼
Patch Embedding
[64, 32, 4, 4]
│
▼
Flatten spatial dimensions
[64, 32, 16]
│
▼
Transpose
[64, 16, 32]
│
▼
Add CLS token
[64, 17, 32]
│
▼
Add positional embedding
[64, 17, 32]
│
▼
Transformer Encoder × 4
[64, 17, 32]
│
▼
Select CLS token
[64, 32]
│
▼
Classification head
[64, 10]
Where:
| Parameter | Value | Meaning |
|---|---|---|
| Image size | 28 × 28 |
MNIST image dimensions |
| Channels | 1 |
Grayscale image |
| Patch size | 7 × 7 |
Size of each image patch |
| Patches | 16 |
(28 / 7)² = 16 |
| Token dimension | 32 |
Embedding size of each token |
| Attention heads | 4 |
Number of attention heads |
| Transformer blocks | 4 |
Number of encoder blocks |
| MLP hidden dimension | 64 |
Hidden size of Transformer MLP |
| Classes | 10 |
Digits 0–9 |
| Batch size | 64 |
Images per training batch |
| Learning rate | 3e-4 |
Adam learning rate |
| Epochs | 5 |
Training epochs |
The notebook is organized around the following components.
The original MNIST image is:
28 × 28
With a patch size of:
7 × 7
we obtain:
28 / 7 = 4 patches per dimension
4 × 4 = 16 patches
The implementation uses:
nn.Conv2d(
in_channels=1,
out_channels=32,
kernel_size=7,
stride=7
)The convolution simultaneously:
- Extracts non-overlapping
7 × 7patches. - Projects each patch into a
32-dimensional embedding.
The resulting tensor is converted from:
[B, 32, 4, 4]
to:
[B, 16, 32]
where B is the batch size.
A learnable classification token is added to the beginning of the patch sequence:
[P1, P2, P3, ..., P16]
becomes:
[CLS, P1, P2, P3, ..., P16]
Therefore:
16 patches + 1 CLS token = 17 tokens
The [CLS] token is eventually used as the representation of the complete image.
Transformers do not inherently know the spatial position of tokens.
Therefore, a learnable positional embedding is added:
Token representation
+
Position representation
↓
Position-aware token representation
The positional embedding has shape:
[1, 17, 32]
because there are 17 tokens and every token has a 32-dimensional representation.
The Transformer uses PyTorch's:
nn.MultiheadAttention(
embed_dim=32,
num_heads=4,
batch_first=True
)The same tensor is supplied as Query, Key and Value:
self.multihead_attention(x, x, x)Therefore:
Q = X
K = X
V = X
This is self-attention.
With:
token_dim = 32
num_heads = 4
each attention head works with:
32 / 4 = 8 dimensions
Conceptually:
32-dimensional token
│
┌───────────────┼───────────────┐
▼ ▼ ▼
Head 1 Head 2 ... Head 4
8-D 8-D 8-D
│ │ │
└───────────────┴───────────────────┘
│
▼
Combined output
32-D
Each TransformerEncoder contains two major sub-blocks.
Input
│
▼
LayerNorm
│
▼
Multi-Head Self-Attention
│
▼
Add Residual
Attention Output
│
▼
LayerNorm
│
▼
Linear: 32 → 64
│
▼
GELU
│
▼
Linear: 64 → 32
│
▼
Add Residual
The complete block is therefore:
Input
│
┌────────┴────────┐
│ │
│ Residual
▼ │
LayerNorm │
│ │
▼ │
Self-Attention │
│ │
└─────── + ◄──────┘
│
┌───────┴────────┐
│ │
│ Residual
▼ │
LayerNorm │
│ │
▼ │
Linear │
32 → 64 │
│ │
▼ │
GELU │
│ │
▼ │
Linear │
64 → 32 │
│ │
└────── + ◄──────┘
│
▼
Output
The notebook uses 4 of these encoder blocks sequentially.
After the Transformer blocks process all 17 tokens, only the first token is selected:
x = x[:, 0]This extracts the [CLS] token:
[B, 17, 32]
│
▼
[B, 32]
The classification head then performs:
LayerNorm
↓
Linear
32 → 10
↓
10 logits
The ten outputs correspond to:
0 1 2 3 4 5 6 7 8 9
The predicted class is obtained using:
outputs.argmax(dim=1)The training process follows the standard deep-learning loop:
MNIST Batch
│
▼
Vision Transformer
│
▼
Predicted Logits
│
▼
CrossEntropyLoss
│
▼
loss.backward()
│
▼
Gradients
│
▼
Adam Optimizer
│
▼
Updated Parameters
The notebook uses:
optimizer = torch.optim.Adam(
model.parameters(),
lr=3e-4
)and:
criterion = nn.CrossEntropyLoss()After training, the model is switched to evaluation mode:
model.eval()and gradient calculation is disabled:
with torch.no_grad():
...The model then predicts the MNIST test set and calculates:
Accuracy =
Correct Predictions / Total Predictions × 100
The notebook also visualizes sample predictions:
┌───────┬───────┬───────┬───────┬───────┐
│ 7 │ 2 │ 1 │ 0 │ 4 │
│Pred:7 │Pred:2 │Pred:1 │Pred:0 │Pred:4 │
│True:7 │True:2 │True:1 │True:0 │True:4 │
└───────┴───────┴───────┴───────┴───────┘
git clone <your-repository-url>
cd vision_transformerpython -m venv .venv
.venv\Scripts\activatepython3 -m venv .venv
source .venv/bin/activateInstall the required packages:
pip install torch torchvision matplotlib jupyterYou can also create a requirements.txt file:
torch
torchvision
matplotlib
jupyter
and install everything with:
pip install -r requirements.txtFrom the project directory:
jupyter notebookor:
jupyter labOpen:
vision_transformer.ipynb
and execute the cells from top to bottom.
If you are using VS Code, install the Python and Jupyter extensions and select the virtual environment as the notebook kernel.
vision_transformer/
│
├── src/
│ └── vision_transformer.ipynb
│
├── data/
│ └── MNIST/
│
├── requirements.txt
│
└── README.md
Note: the notebook currently uses
root="./data". This path is relative to the notebook's current working directory. If you run the notebook fromsrc/, the dataset may be downloaded undersrc/data/. Keep the working directory consistent, or adjust the dataset path if you want the data stored at the project root.
If you're learning Vision Transformers, don't just run the notebook. Try rebuilding it component by component.
A good implementation order is:
Start with:
torchvision.datasets.MNIST(...)Understand the shape of one image:
[1, 28, 28]
Build:
DataLoader(...)and verify:
images.shape
labels.shapeExpected:
images → [64, 1, 28, 28]
labels → [64]
Implement:
class PatchEmbedding(nn.Module):
...Verify that:
[64, 1, 28, 28]
↓
[64, 16, 32]
Build:
class TransformerEncoder(nn.Module):
...Make sure you understand:
- LayerNorm
- Query, Key and Value
- Multi-Head Attention
- Residual connections
- GELU
- Feed-forward MLP
Implement:
class MLPHead(nn.Module):
...and map:
32 → 10
Create:
class VisionTransformer(nn.Module):
...The model should perform:
Image
↓
Patch Embedding
↓
CLS + Position Embedding
↓
Transformer Blocks
↓
CLS Token
↓
Classification Head
↓
10 Classes
Implement the standard:
for epoch in range(epochs):
...training loop.
Use:
model.eval()and:
torch.no_grad()to calculate test accuracy.
Finally, display several MNIST images with:
Predicted label
True label
This makes it easier to see what your model has actually learned.
Once the basic implementation works, modify one variable at a time and observe what happens.
Try changing:
patch_size
token_dim
num_heads
transformer_blocks
mlp_hidden_dim
learning_rate
batch_size
epochsFor example:
Compare:
patch_size = 7
with:
patch_size = 4
Ask yourself:
How does changing the patch size change the number of tokens?
Compare:
num_heads = 1
num_heads = 2
num_heads = 4
Ask:
What happens to the dimensionality of each attention head?
Compare:
transformer_blocks = 1
transformer_blocks = 2
transformer_blocks = 4
Ask:
Does making the Transformer deeper always improve performance?
Try different embedding dimensions while keeping:
token_dim % num_heads == 0
For example:
32 / 4 = 8
64 / 4 = 16
For Multi-Head Attention:
token_dim % num_heads == 0
must hold.
In this project:
32 % 4 = 0
so the configuration is valid.
The model produces logits:
outputs = model(images)and these can be passed directly to:
criterion = nn.CrossEntropyLoss()
loss = criterion(outputs, labels)You do not need to manually apply:
softmax()before CrossEntropyLoss.
If you're new to Transformers, study the project in this order:
1. PyTorch tensors
↓
2. CNN / Conv2d basics
↓
3. Image patches
↓
4. Embeddings
↓
5. Query / Key / Value
↓
6. Self-Attention
↓
7. Multi-Head Attention
↓
8. LayerNorm
↓
9. Residual Connections
↓
10. Transformer Encoder
↓
11. Vision Transformer
The most important thing is to track tensor dimensions throughout the model.
A Transformer attention mechanism can be summarized as:
QKᵀ
Attention = ────── V
√dₖ
where:
Q= QueryK= KeyV= Valuedₖ= dimensionality of the key vectors
The attention mechanism allows one image patch to incorporate information from other patches.
For example:
Patch A ──────┐
Patch B ──────┼──→ Self-Attention ──→ Context-aware representations
Patch C ──────┤
Patch D ──────┘
This is the fundamental idea that allows a Vision Transformer to model relationships between different regions of an image.
The main objective of this project is not simply achieving high MNIST accuracy.
The real objective is to understand:
How can a Transformer designed for sequences process an image?
The answer is:
Image
↓
Divide into patches
↓
Convert patches into embeddings
↓
Treat embeddings as a sequence
↓
Add positional information
↓
Process with Transformer
↓
Use [CLS] representation
↓
Classify the image
Once this pipeline is clear, the architecture of more advanced Vision Transformers becomes much easier to understand.
This implementation intentionally uses standard PyTorch components such as:
nn.Conv2d
nn.LayerNorm
nn.MultiheadAttention
nn.Linear
nn.GELUbut the Vision Transformer architecture itself is assembled manually.
This makes the project useful as a learning exercise before moving to more sophisticated implementations and libraries.
If you're learning Vision Transformers, consider experimenting with the architecture rather than treating the notebook as a black box.
Change one component → train → measure → understand why the result changed.
That's where the real learning begins.