Skip to content

About

A beginner-friendly Vision Transformer (ViT) implemented from scratch using PyTorch and trained on MNIST. Covers patch embeddings, learnable CLS tokens, positional embeddings, multi-head self-attention, Transformer encoder blocks, residual connections, and image classification.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Latest commit

 

History

2 Commits

Folders and files

Repository files navigation

Vision Transformer (ViT) from Scratch with PyTorch

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.


📌 What You'll Learn

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

🧠 Vision Transformer Architecture

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

🔢 Understanding the Tensor Dimensions

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

🧩 Project Components

The notebook is organized around the following components.

1. Patch Embedding

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:

  1. Extracts non-overlapping 7 × 7 patches.
  2. 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.


2. [CLS] Token

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.


3. Positional Embedding

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.


👁️ Multi-Head Self-Attention

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

🔄 Transformer Encoder Block

Each TransformerEncoder contains two major sub-blocks.

Self-Attention Block

Input
  │
  ▼
LayerNorm
  │
  ▼
Multi-Head Self-Attention
  │
  ▼
Add Residual

Feed-Forward Block

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.


🎯 Classification

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)

⚙️ Training Pipeline

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()

🧪 Evaluation

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 │
└───────┴───────┴───────┴───────┴───────┘

🚀 Getting Started

1. Clone the Repository

git clone <your-repository-url>
cd vision_transformer

2. Create a Virtual Environment

Windows

python -m venv .venv
.venv\Scripts\activate

Linux/macOS

python3 -m venv .venv
source .venv/bin/activate

3. Install Dependencies

Install the required packages:

pip install torch torchvision matplotlib jupyter

You can also create a requirements.txt file:

torch
torchvision
matplotlib
jupyter

and install everything with:

pip install -r requirements.txt

▶️ Running the Notebook

From the project directory:

jupyter notebook

or:

jupyter lab

Open:

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.


📁 Recommended Project Structure

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 from src/, the dataset may be downloaded under src/data/. Keep the working directory consistent, or adjust the dataset path if you want the data stored at the project root.


🛠️ Implement It Yourself

If you're learning Vision Transformers, don't just run the notebook. Try rebuilding it component by component.

A good implementation order is:

Step 1 — Load MNIST

Start with:

torchvision.datasets.MNIST(...)

Understand the shape of one image:

[1, 28, 28]

Step 2 — Create DataLoaders

Build:

DataLoader(...)

and verify:

images.shape
labels.shape

Expected:

images → [64, 1, 28, 28]
labels → [64]

Step 3 — Implement Patch Embedding

Implement:

class PatchEmbedding(nn.Module):
    ...

Verify that:

[64, 1, 28, 28]
        ↓
[64, 16, 32]

Step 4 — Implement the Transformer Encoder

Build:

class TransformerEncoder(nn.Module):
    ...

Make sure you understand:

  • LayerNorm
  • Query, Key and Value
  • Multi-Head Attention
  • Residual connections
  • GELU
  • Feed-forward MLP

Step 5 — Build the Classification Head

Implement:

class MLPHead(nn.Module):
    ...

and map:

32 → 10

Step 6 — Combine Everything

Create:

class VisionTransformer(nn.Module):
    ...

The model should perform:

Image
 ↓
Patch Embedding
 ↓
CLS + Position Embedding
 ↓
Transformer Blocks
 ↓
CLS Token
 ↓
Classification Head
 ↓
10 Classes

Step 7 — Train

Implement the standard:

for epoch in range(epochs):
    ...

training loop.

Step 8 — Evaluate

Use:

model.eval()

and:

torch.no_grad()

to calculate test accuracy.

Step 9 — Visualize Predictions

Finally, display several MNIST images with:

Predicted label
True label

This makes it easier to see what your model has actually learned.


💡 Suggested Experiments

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
epochs

For example:

Experiment 1 — Patch Size

Compare:

patch_size = 7

with:

patch_size = 4

Ask yourself:

How does changing the patch size change the number of tokens?

Experiment 2 — Number of Attention Heads

Compare:

num_heads = 1
num_heads = 2
num_heads = 4

Ask:

What happens to the dimensionality of each attention head?

Experiment 3 — Transformer Depth

Compare:

transformer_blocks = 1
transformer_blocks = 2
transformer_blocks = 4

Ask:

Does making the Transformer deeper always improve performance?

Experiment 4 — Token Dimension

Try different embedding dimensions while keeping:

token_dim % num_heads == 0

For example:

32 / 4 = 8
64 / 4 = 16

⚠️ Important Implementation Details

Token Dimension Must Be Divisible by Number of Heads

For Multi-Head Attention:

token_dim % num_heads == 0

must hold.

In this project:

32 % 4 = 0

so the configuration is valid.


Don't Apply Softmax Before CrossEntropyLoss

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.


📚 Recommended Learning Path

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.


🧾 Core Mathematical Idea

A Transformer attention mechanism can be summarized as:

             QKᵀ
Attention = ────── V
             √dₖ

where:

  • Q = Query
  • K = Key
  • V = Value
  • dₖ = 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.


🎓 Learning Objective

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.


🔬 Project Philosophy

This implementation intentionally uses standard PyTorch components such as:

nn.Conv2d
nn.LayerNorm
nn.MultiheadAttention
nn.Linear
nn.GELU

but 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 This Project Helped You

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.

About

A beginner-friendly Vision Transformer (ViT) implemented from scratch using PyTorch and trained on MNIST. Covers patch embeddings, learnable CLS tokens, positional embeddings, multi-head self-attention, Transformer encoder blocks, residual connections, and image classification.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages