Transformer Attention Mechanisms Explained: From Scaled Dot-Product to FlashAttention, GQA, MLA and Sparse Attention

· Deep Learning · By Hassan Nazir

An engineer's guide to how attention works in modern LLMs and why it changed: the scaled dot-product math, KV-cache memory arithmetic, MQA and GQA, DeepSeek's multi-head latent attention, FlashAttention kernels, RoPE context extension, sliding-window and sparse attention, and hybrid linear-attention models.

Every large language model you use is built around one operation: attention. Most of the architectural progress in LLMs from 2023 to 2026 has not replaced it. It has made attention cheaper to store, faster to compute, and able to see further.

This guide walks through the math, then through the engineering changes that matter when you run these models in production.

1. Scaled Dot-Product Attention

The original Transformer paper ("Attention Is All You Need", Vaswani et al., 2017) defines attention as:

Attention(Q, K, V) = softmax( Q·Kᵀ / √d_k ) · V

For each token, the model computes three vectors by multiplying its hidden state with learned matrices:

  • Query (Q): what this token is looking for
  • Key (K): what this token offers to others
  • Value (V): the information this token passes along if selected

The dot product Q·Kᵀ scores how relevant every token is to every other token. Softmax turns scores into weights, and the output is a weighted sum of values.

Why divide by √dk? For random vectors with unit-variance components, the dot product's variance grows with dimension dk. Large logits push softmax into saturated regions with vanishing gradients. Scaling by √d_k keeps the variance near 1.

Causal masking. Decoder-only LLMs add a mask so token i can only attend to tokens at positions ≤ i. This is what makes left-to-right generation possible.

A minimal implementation:

import math
import torch

def attention(q, k, v, causal: bool = True):
    # q, k, v: [batch, heads, seq, head_dim]
    scores = q @ k.transpose(-2, -1) / math.sqrt(q.size(-1))
    if causal:
        seq = q.size(-2)
        mask = torch.triu(torch.ones(seq, seq, dtype=torch.bool, device=q.device), 1)
        scores = scores.masked_fill(mask, float("-inf"))
    return torch.softmax(scores, dim=-1) @ v

In production, use torch.nn.functional.scaleddotproduct_attention, which dispatches to fused kernels such as FlashAttention when the hardware supports them.

2. Multi-Head Attention

Instead of one attention operation over the full hidden dimension, the model splits it into h heads, each with its own projections. Heads learn different relationships (syntax, coreference, positional patterns), and their outputs are concatenated and projected back.

The cost that matters is quadratic in sequence length: the score matrix is n × n per head. Doubling context length quadruples attention compute for that layer during prefill.

3. The KV Cache: Where Inference Memory Goes

During generation, the model produces one token at a time. Recomputing keys and values for all previous tokens at every step would be wasteful, so inference engines cache K and V for every layer.

The KV-cache size per token is:

bytes_per_token = 2 (K and V) × n_layers × n_kv_heads × head_dim × bytes_per_value

Worked example for a Llama-3-70B-style configuration (80 layers, 8 KV heads, head dimension 128, 16-bit values):

2 × 80 × 8 × 128 × 2 bytes = 327,680 bytes ≈ 320 KiB per token
× 131,072 tokens (128K context) ≈ 40 GiB for ONE sequence

That is why long-context serving is memory-bound, and why most attention innovations target the nkvheads × head_dim term.

PagedAttention (introduced with vLLM, Kwon et al., 2023) manages this memory like virtual memory pages, which removes fragmentation and enables sharing prefixes across requests. It is the reason modern serving engines can batch many long sequences.

4. Shrinking the Cache: MQA and GQA

Multi-Query Attention (MQA) (Shazeer, 2019) keeps many query heads but uses a single shared key/value head. The KV cache shrinks by the number of heads, at some cost in quality.

Grouped-Query Attention (GQA) (Ainslie et al., 2023) is the compromise most open models adopted: query heads are divided into groups, and each group shares one KV head.

VariantQuery headsKV headsRelative KV cache
Multi-Head (MHA)64641×
Grouped-Query (GQA)6481/8×
Multi-Query (MQA)6411/64×

GQA recovers most of MHA's quality with a fraction of the memory, which is why it appears in Llama 3, Mistral, Qwen, and many other open-weight families.

def repeat_kv(k, v, n_rep: int):
    # Expand grouped KV heads to match query heads: [b, kv_heads, s, d] -> [b, kv_heads*n_rep, s, d]
    if n_rep == 1:
        return k, v
    b, h, s, d = k.shape
    k = k[:, :, None].expand(b, h, n_rep, s, d).reshape(b, h * n_rep, s, d)
    v = v[:, :, None].expand(b, h, n_rep, s, d).reshape(b, h * n_rep, s, d)
    return k, v

5. Multi-Head Latent Attention (MLA)

DeepSeek-V2 (2024) introduced Multi-Head Latent Attention. Instead of caching full keys and values, MLA caches a compressed low-rank latent vector per token and reconstructs per-head keys and values from it through learned up-projections.

The result is a KV cache far smaller than standard MHA while retaining strong quality. DeepSeek-V3 and later DeepSeek models kept MLA, and it has influenced other architectures. A detail that matters in practice: rotary position embeddings do not commute cleanly with the compression, so MLA uses a small decoupled RoPE component carried alongside the latent.

6. FlashAttention: Same Math, Better Memory Traffic

Attention on GPUs is often limited by memory bandwidth, not arithmetic. A naive implementation writes the full n × n score matrix to high-bandwidth memory (HBM) and reads it back.

FlashAttention (Dao et al., 2022) computes exact attention in tiles that fit in fast on-chip SRAM, using an online softmax so the full matrix is never materialized.

  • FlashAttention-2 (2023) improved parallelism and work partitioning.
  • FlashAttention-3 (2024) targets NVIDIA Hopper GPUs, using asynchronous execution and FP8 support.

The key point: FlashAttention changes speed and memory use, not the output. It is an exact algorithm, not an approximation.

7. Position: RoPE and Context Extension

Attention by itself has no sense of order. Most modern LLMs use Rotary Position Embedding (RoPE) (Su et al., 2021), which rotates query and key vectors by an angle proportional to position. The dot product then depends on relative distance between tokens.

Extending a model beyond its training context requires adjusting those rotations:

  • Position Interpolation scales positions down to fit the trained range.
  • NTK-aware scaling changes the RoPE base frequency.
  • YaRN (Peng et al., 2023) combines frequency-dependent interpolation with an attention temperature adjustment and is widely used for long-context extension.

8. Local and Sparse Attention

Full attention over a million tokens is expensive even with FlashAttention. Several techniques restrict which tokens attend to which.

Sliding-window attention. Each token attends only to the previous W tokens. Mistral 7B (2023) popularized it, and several model families interleave local sliding-window layers with periodic global layers to cut KV memory while preserving long-range access.

Attention sinks. StreamingLLM (Xiao et al., 2023) observed that models dump large attention mass on the first few tokens. Keeping those "sink" tokens plus a recent window lets models stream far beyond their window without collapsing.

Learned sparse attention. DeepSeek's Native Sparse Attention paper (2025) combined compressed, selected, and sliding-window branches that are trainable end to end. DeepSeek later shipped DeepSeek Sparse Attention in its V3.2 line, where a lightweight indexer selects which past tokens each query attends to, reducing long-context cost.

9. Beyond Softmax: Linear Attention and Hybrids

State-space and linear-attention models replace the quadratic score matrix with a recurrent state that has constant size per layer:

  • Mamba (Gu and Dao, 2023) introduced selective state-space layers with linear-time sequence processing.
  • Hybrid architectures interleave a minority of full-attention layers with many linear-time layers. Examples include AI21's Jamba (Mamba plus attention) and Qwen3-Next (2025), which mixes Gated DeltaNet layers with gated attention.

The trade-off: linear layers are cheaper for long sequences, but exact recall of specific earlier tokens is harder. Keeping some full-attention layers preserves that capability.

10. What This Means for Engineers Shipping LLM Systems

If you are...Pay attention to...
Sizing GPUs for inferenceKV heads, head dim, layers, context length, and batch size (the KV formula above)
Serving long contextGQA/MLA models, paged KV cache, prefix caching, and quantized KV (FP8)
Building RAGWhether the model actually uses long context well; test retrieval at depth, not just window size
Fine-tuningKeep RoPE scaling settings consistent between training and serving
Optimizing latencyPrefill is compute-bound, decode is memory-bound; tune them separately

Key Takeaways

  • Attention is softmax(QKᵀ/√d)V; the cost is quadratic in sequence length and the KV cache dominates inference memory.
  • GQA and MLA shrink the KV cache; FlashAttention reduces memory traffic without changing results.
  • RoPE scaling (PI, NTK, YaRN) enables longer contexts; sliding-window, sink, and sparse attention cut the cost of using them.
  • Hybrid linear-attention models are the main architectural bet for very long contexts, with full-attention layers kept for precise recall.

References

Related field notes