How the Attention Mechanism Works in Transformer Models: From Theory to Implementation

The attention mechanism enables transformer models to dynamically weigh the importance of all other tokens when processing any given token, computing context-aware representations through scaled dot-product attention where queries, keys, and values interact via the formula $\text{Attention}(Q,K,V)=\text{softmax}(QK^\top/\sqrt{d_k})V$.

According to the HenryNdubuaku/maths-cs-ai-compendium repository, this mechanism serves as the computational backbone of modern natural language processing, allowing parallel processing of sequences while capturing long-range dependencies that traditional recurrent architectures struggle to model.

The Mathematical Foundation of Self-Attention

Scaled Dot-Product Attention

At its core, the attention mechanism computes a weighted average of values based on the similarity between queries and keys. The operation follows the scaled dot-product attention formula:

$$ \text{Attention}(Q,K,V)=\text{softmax}!\left(\frac{QK^{\top}}{\sqrt{d_k}}\right)V $$

The scaling factor $\sqrt{d_k}$ (where $d_k$ represents the dimension of the key vectors) prevents the dot products from growing too large in magnitude, which would push the softmax function into regions with extremely small gradients. Before positional encodings are injected, self-attention exhibits permutation equivariance, producing identical outputs regardless of token order, making it essential to add positional information through learned or sinusoidal embeddings.

Query, Key, and Value Projections

In practice, input embeddings undergo linear transformations to generate three distinct representations:

  • Queries (Q): Each token generates a query vector that represents what it is looking for
  • Keys (K): Each token presents a key vector that represents what it contains
  • Values (V): Each token offers a value vector containing the actual content to be aggregated

Raw attention scores emerge from the dot product $QK^\top$, measuring how relevant each key is to each query. After scaling and softmax normalization (ensuring weights sum to 1), these weights combine the value vectors, yielding a context-aware representation for each position.

Multi-Head Attention Architecture

Transformers employ multi-head attention to capture diverse relational patterns simultaneously. Rather than computing attention once, the model runs $h$ independent attention heads in parallel, each with its own projection matrices $(W_Q^i, W_K^i, W_V^i)$.

The outputs from all heads are concatenated and projected through a final linear layer, allowing the model to attend to information from different representation subspaces at different positions—simultaneously capturing syntactic relationships, semantic associations, and long-range dependencies.

Efficient Attention Variants for Long Sequences

Standard attention scales quadratically with sequence length ($O(n^2)$), creating memory and compute bottlenecks for long contexts. The compendium documents several architectural refinements that modify the attention mechanism for efficiency.

Full Quadratic Attention and Its Limitations

Full (quadratic) attention—where every token attends to all others—remains the standard in early transformers like BERT and GPT. While maximally expressive, this approach becomes prohibitively expensive for long documents or high-resolution inputs, as detailed in chapter 17 - AI inference/02. efficient architectures.md.

Sliding-Window and Sparse Patterns

Sliding-window attention restricts each token to attend only to a fixed-size recent window, reducing complexity to $O(n \cdot w)$ where $w$ is the window size. This pattern, used in models like Mistral and Gemma, enables processing of long contexts with modest memory requirements.

Local + global attention implements a hybrid approach where most tokens use a restricted window while designated "global" tokens attend to the entire sequence. This architecture, found in Longformer and BigBird, balances local pattern recognition with long-range dependency modeling. Sparse/dilated attention further reduces computation by attending to every $k$-th token within a window, creating hierarchical coverage patterns.

Flash Attention and Memory Efficiency

Flash Attention revolutionized attention computation by reformulating the algorithm to exploit memory hierarchy rather than approximating the attention matrix. As documented in chapter 16 - SIMD and GPU programming/05. triton, TPUs and pallax.md, Flash Attention tiles the $QK^\top$ multiplication into small blocks computed in fast on-chip SRAM, achieving $O(n)$ memory usage while producing exact (not approximate) softmax results. This tiling strategy provides 2–4× speedups and has become the default in major frameworks.

Linear and Multi-Query Attention

Linear attention replaces the softmax with kernel feature approximations, reducing time complexity to $O(n)$. While faster for very long sequences, these methods often trade off expressiveness compared to full attention.

Multi-Query Attention (MQA) and Grouped-Query Attention (GQA) optimize inference by sharing a single key-value cache across attention heads, dramatically reducing the KV-cache size. This optimization proves critical for efficient serving on GPUs and TPUs, as discussed in chapter 17 - AI inference/03. serving and batching.md, which also covers PagedAttention for managing attention key-values in paged memory blocks.

Attention Beyond Text

The same attention formulation extends beyond language to graph structured data. chapter 12 - graph neural networks/04. graph attention networks.md demonstrates how the query-key-value mechanism adapts to irregular graph topologies, allowing nodes to attend selectively to their neighbors.

Implementing Attention in PyTorch

Below are minimal, framework-agnostic implementations illustrating the core attention mechanism and its sliding-window variant.

import torch
import torch.nn.functional as F

def scaled_dot_product_attention(Q, K, V, mask=None):
    """
    Q, K, V: tensors of shape (batch, heads, seq_len, head_dim)
    mask: optional bool tensor broadcasting to (batch, heads, seq_len, seq_len)
    """
    d_k = Q.size(-1)
    scores = torch.matmul(Q, K.transpose(-2, -1)) / d_k**0.5   # (B, H, L, L)

    if mask is not None:
        scores = scores.masked_fill(~mask, float('-inf'))

    attn_weights = F.softmax(scores, dim=-1)                 # (B, H, L, L)

    output = torch.matmul(attn_weights, V)                  # (B, H, L, head_dim)

    return output, attn_weights

# ----- Example: Sliding-window attention (window size = 4) -----

def sliding_window_attention(Q, K, V, window=4):
    B, H, L, D = Q.shape
    output = torch.zeros_like(Q)
    attn = torch.zeros(B, H, L, L, device=Q.device)

    for i in range(L):
        start = max(0, i - window + 1)
        q = Q[:, :, i:i+1, :]                     # (B, H, 1, D)

        k = K[:, :, start:i+1, :]                 # (B, H, ≤window, D)

        v = V[:, :, start:i+1, :]                 # (B, H, ≤window, D)

        o, w = scaled_dot_product_attention(q, k, v)   # → (B, H, 1, D)

        output[:, :, i, :] = o.squeeze(2)
        attn[:, :, i, start:i+1] = w.squeeze(2)
    return output, attn

The first function implements the textbook scaled dot-product attention found in the original transformer paper. The second demonstrates a naive sliding-window scheme with $O(L \cdot w)$ complexity, illustrating how attention patterns can be restricted to improve efficiency.

Summary

  • Scaled dot-product attention forms the foundation of transformer models, computing token interactions as $\text{softmax}(QK^\top/\sqrt{d_k})V$ to produce context-aware representations.
  • Multi-head attention parallelizes multiple attention operations, enabling the model to capture diverse syntactic and semantic patterns simultaneously.
  • Memory bandwidth limitations often constrain attention more than raw compute, motivating optimizations like Flash Attention and KV-cache reduction techniques.
  • Sparse attention variants—including sliding-window, local+global, and dilated patterns—reduce complexity from $O(n^2)$ to linear or near-linear while preserving model expressiveness.
  • Flash Attention achieves exact attention computation with $O(n)$ memory through careful tiling and SRAM management, as implemented in chapter 16 - SIMD and GPU programming/05. triton, TPUs and pallax.md.

Frequently Asked Questions

Why is the attention score divided by $\sqrt{d_k}$?

The scaling factor $\sqrt{d_k}$ prevents the dot products $QK^\top$ from growing too large as dimensionality increases. Without this scaling, the softmax function would receive extremely large values, producing vanishing gradients and unstable training. Dividing by the square root of the key dimension keeps the variance of the dot products roughly constant, ensuring the softmax operates in a well-behaved range.

How does Flash Attention achieve linear memory complexity?

Flash Attention reformulates the attention computation to avoid materializing the full $n \times n$ attention matrix in high-bandwidth memory. Instead, it tiles the computation into small blocks that fit in fast on-chip SRAM, computing the softmax incrementally using online statistics. This approach performs exact attention (unlike sparse approximations) while requiring only $O(n)$ memory rather than $O(n^2)$, as detailed in chapter 16 - SIMD and GPU programming/05. triton, TPUs and pallax.md.

What is the difference between multi-head attention and multi-query attention?

Multi-head attention uses separate key and value projections for each attention head, requiring $h \times$ more memory for the KV-cache during inference. Multi-query attention shares a single key and value projection across all heads, reducing the KV-cache size by the number of heads while maintaining separate query projections. This optimization dramatically improves inference throughput for long sequences with minimal quality degradation.

Why do transformers need positional encodings if attention already captures relationships?

Self-attention is fundamentally permutation equivariant, meaning it treats the input as an unordered set and produces identical outputs regardless of token order. Since language and sequential data depend crucially on order, transformers inject positional information through either learned embeddings or sinusoidal encodings added to the input embeddings. These encodings break the permutation symmetry, allowing the model to distinguish between different positions in the sequence.

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 →