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

> Learn to implement multi-head attention from scratch using PyTorch. Understand how parallel heads and scaled dot-product attention work in this detailed guide. Perfect for LLM development.

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

---

**Multi-head attention splits the embedding dimension into parallel heads that each compute scaled dot-product attention with causal masking; the `train-llm-from-scratch` repository implements this via a `Head` class for single-head logic and a `MultiHeadAttention` class that concatenates multiple heads along the embedding dimension to preserve the input shape `(B, T, n_embed)`.**

Multi-head attention is the core mechanism powering modern Transformer architectures. This article examines a minimal PyTorch implementation from the `FareedKhan-dev/train-llm-from-scratch` repository, demonstrating exactly how to implement multi-head attention from scratch using PyTorch without relying on high-level abstractions.

## Understanding the Multi-Head Attention Architecture

The implementation follows the original "Attention Is All You Need" design, separating concerns into two distinct classes within [`src/models/attention.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py). The `Head` class manages single-head attention logic, while `MultiHeadAttention` orchestrates parallel computation across multiple heads.

This modular approach allows each head to learn different representation subspaces while maintaining a clean interface for the broader Transformer model.

## Step-by-Step Implementation

### Single Attention Head (`Head` class)

Located at lines 6-34 in [`src/models/attention.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py), the `Head` class implements scaled dot-product attention with three critical operations:

**Linear projections.** For each input tensor of shape `(B, T, n_embed)`, three `nn.Linear` layers—`key`, `query`, and `value`—project the embedding into a lower-dimensional `head_size`. The implementation omits biases for computational efficiency, mapping to `head_size = n_embed // n_head`.

**Scaled dot-product computation.** The attention scores are computed as the matrix multiplication of queries (`q`) and transposed keys (`k.transpose(-2, -1)`), producing a raw score matrix of shape `(B, T, T)`. These scores are scaled by `1 / sqrt(head_size)` (the `scale_factor`) to prevent softmax saturation and stabilize gradients.

**Causal masking.** A lower-triangular matrix registered as buffer `self.tril` masks future positions with `float('-inf')`, enforcing autoregressive behavior where each token attends only to previous tokens.

The final output results from applying softmax to the masked scores and multiplying by the values tensor (`v`).

### Parallel Head Aggregation (`MultiHeadAttention` class)

The `MultiHeadAttention` class (lines 60-96 in [`src/models/attention.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py)) aggregates multiple `Head` instances:

1. **ModuleList initialization.** The constructor creates `n_head` instances of `Head`, each processing a slice of the embedding dimension in parallel.

2. **Concatenation strategy.** During the forward pass, each head produces an output of shape `(B, T, head_size)`. The implementation concatenates these along the last dimension (`dim=-1`), yielding a tensor of shape `(B, T, n_embed)` that matches the input dimensions but contains richer contextual information.

This shape preservation allows the attention output to be directly added to the residual stream within Transformer blocks.

## Complete Code Examples

### Stand-Alone Multi-Head Attention

The following example mirrors the test block at lines 98-111 in [`src/models/attention.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py), demonstrating initialization and forward pass:

```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

# 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 preserves input shape

out = mh_attn(x)
print(f"Input shape:  {x.shape}")   # torch.Size([2, 5, 32])

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

```

### Integration into a Transformer Block

In practice, the attention module integrates into a full Transformer block via [`src/models/transformer_block.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/transformer_block.py) (lines 6-30), which wraps the attention with layer normalization and residual connections:

```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 (B, T, embed_dim)

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

```

## Key Implementation Details

Several design choices in `FareedKhan-dev/train-llm-from-scratch` optimize the implementation for clarity and training stability:

- **Head dimension calculation.** The `head_size` is strictly derived as `n_embed // n_head`, ensuring the concatenated heads exactly reconstruct the original embedding dimension without additional projection layers.
  
- **Buffer-based masking.** The causal mask `self.tril` is registered as a persistent buffer rather than computed dynamically, reducing overhead during training loops.

- **Biases disabled.** All `nn.Linear` projections for keys, queries, and values set `bias=False`, following modern LLM implementations to reduce parameter count without impacting representational capacity.

## Summary

- The `Head` class in [`src/models/attention.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py) implements single-head scaled dot-product attention with causal masking via a lower-triangular buffer.
- `MultiHeadAttention` parallelizes `n_head` instances using `nn.ModuleList` and concatenates results along `dim=-1` to preserve the `(B, T, n_embed)` tensor shape.
- Scaled attention (`1/sqrt(head_size)`) prevents softmax gradient vanishing in deep models.
- The implementation integrates into full Transformer blocks via [`src/models/transformer_block.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/transformer_block.py), supporting autoregressive language modeling workflows.

## Frequently Asked Questions

### Why is the attention score scaled by the square root of the head size?

Scaling by `1 / sqrt(head_size)` prevents the dot-product values from growing too large in magnitude as the dimensionality increases. Without this scaling, the softmax function would produce extremely small gradients, slowing or destabilizing training. This scaling factor appears in lines 6-34 of [`src/models/attention.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py).

### What is the purpose of the causal mask in the `Head` class?

The causal mask ensures autoregressive behavior by preventing tokens from attending to future positions in the sequence. Implemented via the `self.tril` buffer in [`src/models/attention.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py), it masks the upper triangle of the attention score matrix with `float('-inf')` before the softmax operation, effectively setting future attention weights to zero.

### How does `MultiHeadAttention` maintain the input-output shape consistency?

The class divides the embedding dimension evenly across heads (`head_size = n_embed // n_head`), processes each head in parallel, then concatenates the results along the last dimension. Because `n_head * head_size = n_embed`, the output tensor retains the original shape `(B, T, n_embed)`, allowing seamless residual connections in the Transformer architecture.

### Can this implementation handle variable sequence lengths?

Yes. The implementation uses PyTorch tensor operations that support dynamic shapes. While the `context_length` parameter initializes the causal mask buffer, the forward pass in [`src/models/attention.py`](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py) handles any sequence length up to this maximum, making it suitable for training and inference with varying text lengths.