# How to Implement Multi-Head Attention from Scratch in PyTorch

> Implement multi-head attention from scratch in PyTorch. Learn to build scaled dot-product attention with causal masking and combine multiple heads for richer contextual representation.

- Repository: [Fareed Khan/train-llm-from-scratch](https://github.com/FareedKhan-dev/train-llm-from-scratch)
- Tags: how-to-guide
- Published: 2026-05-31

---

**Implement multi-head attention in PyTorch by creating a `Head` class for scaled dot-product attention with causal masking, then wrapping multiple heads in a `MultiHeadAttention` class that concatenates their outputs to preserve input dimensions while enriching contextual representation.**

To implement multi-head attention from scratch in PyTorch, you need to split the embedding dimension into parallel attention heads, compute scaled dot-product attention for each, and concatenate the results. The repository FareedKhan-dev/train-llm-from-scratch provides a minimal, educational implementation following the original "Attention Is All You Need" paper. The code demonstrates how to build the core Transformer mechanism without relying on high-level abstractions like `nn.MultiheadAttention`.

## Understanding the Architecture

The implementation splits attention into two distinct classes to separate single-head logic from multi-head orchestration.

### The Single Attention Head

The `Head` class in [`src/models/attention.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py) (lines 6-34) implements one attention head. It projects the input into **queries**, **keys**, and **values** using three separate `nn.Linear` layers. Each projection maps the full embedding dimension `n_embed` down to `head_size = n_embed // n_head`. Critically, these linear layers omit biases for simplicity and computational efficiency.

### Combining Multiple Heads

The `MultiHeadAttention` class in [`src/models/attention.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py) (lines 60-96) manages parallelism. It initializes a `nn.ModuleList` of `Head` instances according to the `n_head` parameter. During the forward pass, it processes each head independently, concatenates their outputs along the embedding dimension (`dim=-1`), and returns a tensor with shape `(B, T, n_embed)` identical to the input, but containing richer contextual information from multiple representation subspaces.

## Step-by-Step Implementation Details

The attention mechanism follows five precise operations as implemented in the source code:

1. **Linear Projections** – Three `nn.Linear` layers (`key`, `query`, `value`) transform the input from `n_embed` to `head_size` dimensions.

2. **Scaled Dot-Product** – Compute raw attention scores by multiplying queries (`q`) with the transpose of keys (`k.T`), yielding a matrix of shape `(B, T, T)`. Scale these scores by `1/√head_size` (stored as `scale_factor`) to prevent softmax saturation and maintain gradient stability.

3. **Causal Masking** – Register a lower-triangular matrix (`self.tril`) as a buffer during initialization. Apply this mask to the attention scores by setting future positions to `float('-inf')`, ensuring tokens only attend to previous positions for autoregressive generation.

4. **Softmax and Weighted Sum** – Apply softmax over the last dimension to obtain attention weights, then multiply these weights with the values (`v`) to produce the head output.

5. **Head Concatenation** – Concatenate outputs from all heads along the embedding dimension, resulting in a tensor that maintains the original input shape but combines information from all attention heads.

## Code Examples

### Stand-Alone Multi-Head Attention

Use the `MultiHeadAttention` class directly for testing or integration into custom architectures. This example mirrors the validation block found at lines 98-111 of [`src/models/attention.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py):

```python
import torch
from src.models.attention import MultiHeadAttention

# Hyper-parameters

batch_size  = 2
seq_len     = 5
embed_dim   = 32
num_heads   = 4
context_len = seq_len

# Random input tensor (Batch, Time, Channels)

x = torch.randn(batch_size, seq_len, embed_dim)

# Initialize multi-head attention

mh_attn = MultiHeadAttention(
    n_head=num_heads,
    n_embed=embed_dim,
    context_length=context_len,
)

# Forward pass

out = mh_attn(x)

print("Input shape :", x.shape)   # → torch.Size([2, 5, 32])

print("Output shape:", out.shape)  # → torch.Size([2, 5, 32])

```

### Integration into Transformer Blocks

In practice, multi-head attention operates within a larger Transformer block. The `Block` class in [`src/models/transformer_block.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/transformer_block.py) (lines 6-30) wraps the attention mechanism with layer normalization, residual connections, and a feedforward MLP:

```python
from src.models.transformer_block import Block

block = Block(
    n_head=num_heads,
    n_embed=embed_dim,
    context_length=context_len,
)

output = block(x)  # Shape: (Batch, Time, embed_dim)

print(output.shape)  # → torch.Size([2, 5, 32])

```

## Source Code Structure

Reference these specific files in FareedKhan-dev/train-llm-from-scratch to explore the complete implementation:

- **[`src/models/attention.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py)** – Contains the `Head` and `MultiHeadAttention` class definitions implementing the core attention logic.

- **[`src/models/transformer_block.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/transformer_block.py)** – Demonstrates how attention integrates with layer normalization, residual connections, and MLP layers to form a complete Transformer block.

- **[`src/models/transformer.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/transformer.py)** – Stacks multiple `Block` instances to construct a full Transformer language model.

- **[`scripts/train_transformer.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/scripts/train_transformer.py)** – Provides the training script that wires the model, dataset, and optimizer together for end-to-end training.

## Summary

- **Split embeddings across heads** by setting `head_size = n_embed // n_head` to enable parallel attention processing.
- **Scale attention scores** by `1/√head_size` before softmax to maintain stable gradients during training.
- **Apply causal masking** using a lower-triangular buffer to enforce autoregressive behavior in decoder-only models.
- **Concatenate head outputs** to preserve the original tensor shape while combining contextual information from multiple representation subspaces.
- **Omit biases** in projection layers for computational efficiency, following modern Transformer implementations.

## Frequently Asked Questions

### What is the difference between the Head and MultiHeadAttention classes?

The `Head` class implements a single attention mechanism with its own query, key, and value projections, computing scaled dot-product attention for one representation subspace. The `MultiHeadAttention` class instantiates multiple `Head` objects in parallel, processes the input through each independently, and concatenates their outputs to combine multiple perspectives simultaneously.

### Why must attention scores be scaled by 1 divided by the square root of head size?

Scaling by `1/√head_size` (also called the scale factor) prevents the dot-product values from growing too large as dimensionality increases. According to the Transformer paper, large dot-products push the softmax function into regions with extremely small gradients, which slows training convergence. The scaling factor keeps the variance of attention weights manageable and preserves gradient flow.

### How does causal masking work in this implementation?

The implementation registers a lower-triangular boolean matrix (`self.tril`) as a buffer during initialization. During the forward pass, it masks the attention scores by setting future positions (above the diagonal) to negative infinity before the softmax operation. This ensures each position can only attend to itself and previous positions, which is essential for language modeling tasks where future tokens must remain hidden.

### Can this implementation support bidirectional (non-causal) attention?

Yes, you can modify the implementation for bidirectional attention by removing or conditionally applying the causal mask. In [`src/models/attention.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py), you would either skip the masking step where `self.tril` is applied, or add a boolean parameter (e.g., `causal=True`) that controls whether to apply the mask. This flexibility allows the same code to support both decoder-only (causal) and encoder-style (bidirectional) Transformer architectures.