How to Implement Self-Attention Mechanism from Scratch in Python

Self-attention enables transformer models to relate every token in a sequence to every other token by computing scaled dot-product attention between learned query, key, and value projections, implemented efficiently in PyTorch using a single combined QKV linear layer and causal masking for autoregressive generation.

Self-attention forms the computational backbone of modern large language models, allowing sequences to capture long-range dependencies without recurrence. The rohitg00/ai-engineering-from-scratch repository provides a complete, educational implementation that demonstrates exactly how to build this mechanism using pure PyTorch operations. This guide walks through the actual source code from the capstone projects, covering everything from linear projections to multi-head processing.

The Seven Steps of Self-Attention

The MultiHeadSelfAttention class in phases/19-capstone-projects/33-multihead-self-attention/code/main.py encapsulates the complete attention mechanism through seven distinct operations:

  1. Linear projections – A single learned layer generates queries (Q), keys (K), and values (V) from the input representation x with shape B×T×D
  2. Head splitting – The concatenated QKV tensors reshape so each head processes a sub-space of dimension d_head = D / n_heads
  3. Scaled dot-product – Attention scores compute as Q·Kᵀ scaled by √d_head to maintain stable softmax gradients
  4. Causal mask – A lower-triangular mask prevents tokens from attending to future positions in autoregressive setups
  5. Softmax and dropout – Masked scores convert to probabilities via softmax, followed by optional regularization
  6. Weighted sum – The probability distribution combines values through matrix multiplication weights·V
  7. Head merging and output projection – Per-head results concatenate and project back to the original dimension

Building the MultiHeadSelfAttention Class

Combined QKV Projection

For computational efficiency, the implementation uses a single linear layer to project the input into query, key, and value representations simultaneously. In phases/19-capstone-projects/33-multihead-self-attention/code/main.py at line 47, the initialization creates:

self.qkv_proj = nn.Linear(d_model, 3 * d_model, bias=True)

This concatenated projection reduces memory overhead and kernel launch costs compared to three separate layers. The forward pass splits the output tensor into three equal parts corresponding to Q, K, and V.

Splitting and Merging Heads

Multi-head attention divides the model dimension across parallel attention heads. The _split_heads method reshapes the tensor using view and transpose operations:

x.view(b, t, n_heads, d_head).transpose(1, 2)

This transforms the shape from (batch, seq, d_model) to (batch, n_heads, seq, d_head), allowing each head to attend to the full sequence independently. After computing attention, the _merge_heads method reverses this operation at line 60:

x.transpose(1, 2).contiguous().view(b, t, h * dh)

Scaled Dot-Product Attention

The core attention computation occurs at line 83, where the implementation calculates compatibility scores between queries and keys:

scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_head)

Scaling by the square root of the head dimension prevents the dot products from growing too large in magnitude, which would push the softmax function into regions with extremely small gradients.

Causal Masking Implementation

For language modeling tasks, the class maintains a causal mask buffer initialized at lines 52-53:

self.causal_mask = torch.tril(torch.ones(max_context_length, max_context_length))

During the forward pass, the implementation slices this mask to the current sequence length and applies it using masked_fill:

scores = scores.masked_fill(mask_slice == 0, float("-inf"))

This ensures positions can only attend to themselves and previous tokens, maintaining the autoregressive property required for next-token prediction.

Practical Implementation Examples

Minimal Self-Attention Demo

To use the implementation directly for experimentation:

import torch
from phases.19_capstone_projects.33_multihead_self_attention.code.main import MultiHeadSelfAttention

# Configuration

batch, seq_len, d_model = 2, 5, 16
n_heads = 4

# Random input

x = torch.randn(batch, seq_len, d_model)

# Initialise the layer

attn = MultiHeadSelfAttention(d_model=d_model,
                              n_heads=n_heads,
                              max_context_length=seq_len)

# Forward pass – returns (output, attention_weights)

out, weights = attn(x, return_weights=True)

print("output shape :", out.shape)           # → (2, 5, 16)

print("weights shape:", weights.shape)      # → (2, 4, 5, 5)

Tiny Language Model Integration

The repository includes TinyAttentionLM, which demonstrates how self-attention functions within a complete architecture:

from phases.19_capstone_projects.33_multihead_self_attention.code.main import TinyAttentionLM, DemoConfig

cfg = DemoConfig(vocab_size=64, d_model=32, n_heads=4, seq_len=12, batch_size=2)
model = TinyAttentionLM(vocab_size=cfg.vocab_size,
                       d_model=cfg.d_model,
                       n_heads=cfg.n_heads,
                       max_context_length=cfg.seq_len)

# Dummy token ids

ids = torch.randint(0, cfg.vocab_size, (cfg.batch_size, cfg.seq_len))

# Forward pass – get logits and per‑head attention weights

logits, attn_weights = model(ids, return_weights=True)

print("logits shape :", logits.shape)          # → (2, 12, 64)

print("attention weights shape :", attn_weights.shape)  # → (2, 4, 12, 12)

Full Transformer Block Stack

For production-like usage, the attention mechanism integrates with LayerNorm and feed-forward networks in phases/19-capstone-projects/34-transformer-block/code/main.py:

from phases.19_capstone_projects.34_transformer_block.code.main import BlockConfig, BlockStack

cfg = BlockConfig(d_model=128, num_heads=8, context_length=32, pre_ln=True)
stack = BlockStack(cfg, depth=3)   # three transformer blocks

tokens = torch.randint(0, 128, (2, 32))   # batch of token ids

output = stack(tokens)

print("stack output shape:", output.shape)   # → (2, 32, 128)

Summary

  • Single QKV projection improves efficiency by computing queries, keys, and values in one matrix multiplication using nn.Linear(d_model, 3 * d_model)
  • Head splitting via view and transpose operations enables parallel attention across multiple representation subspaces without increasing computational complexity
  • Causal masking uses a pre-allocated lower-triangular buffer to prevent attention to future tokens during training
  • Scaling factor of 1/√d_head maintains stable gradient flow through the softmax operation
  • Modular design allows the MultiHeadSelfAttention class to function standalone or within complete transformer blocks containing residual connections and layer normalization

Frequently Asked Questions

Why combine QKV into a single linear projection?

Combining the query, key, and value projections into one layer reduces memory overhead and enables optimized kernel fusion during training. The implementation splits the output tensor into three parts, achieving identical mathematical results to separate layers while improving computational efficiency on modern hardware accelerators.

What is the purpose of scaling by the square root of the head dimension?

Scaling by √d_head prevents the dot-product values from becoming too large as the dimension increases. Without this scaling, the softmax function would receive extremely large inputs, producing gradients near zero that slow down or prevent effective learning.

How does causal masking prevent looking at future tokens?

The causal mask is a lower-triangular matrix where positions (i, j) contain 1 if j ≤ i and 0 otherwise. Before applying softmax, the implementation sets masked positions to negative infinity using masked_fill, causing the softmax output to be zero for future positions. This ensures each token can only aggregate information from itself and previous tokens in the sequence.

Can this implementation handle bidirectional attention?

Yes, by setting the causal mask parameter to None or modifying the mask initialization logic, the same MultiHeadSelfAttention class supports bidirectional attention for encoder-only architectures like BERT. The mask buffer is optional, and the forward pass can skip the masking step when processing non-autoregressive tasks.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →