Skip to content
AI-grafen
EUniversityTransformer architecture· about 60 min· evolving, reviewed regularly· verified 2026-09-20· EN

Context length and the quadratic cost

Be able to compute the memory and time cost of attention as a function of the sequence length.

Prerequisites

Intuition

Attention compares every token with every other one. With T tokens that is T² pairs — double the context and the cost in the attention step quadruples.

ContextPairsRelative
2 k4 M1×
8 k67 M16×
32 k1.0 G256×
128 k16 G4 096×

But two costs have to be kept apart:

  • Computation in the prefill: O(T²·d) — it grows quadratically.
  • Memory during generation: the KV cache grows linearly in T, but still becomes the bottleneck in practice since it stays there for the whole session.

FlashAttention removed the quadratic memory cost for the attention matrix, but not the quadratic computation.

Formal

FLOPs per layer (approximately, the forward pass):

  • The Q, K, V, O projections: 4⋅Td24\cdot T d^2 — linear in T.
  • The attention scores QK⊤QK^\top and softmax⋅V\text{softmax}\cdot V: 2T2d2T^2 d — quadratic.
  • The MLP (4× expansion): 8Td28Td^2 — linear.

Attention only dominates once T≳dT \gtrsim d. For d=4096d = 4096 that means contexts below about 4 k tokens are dominated by the linear terms — which surprises many people. The quadratic term only bites in earnest at long contexts.

Memory:

  • The KV cache: 2LHkvdheadT2 L H_{kv} d_{head} T per sequence — linear in T, but multiplied by the batch.
  • The attention matrix: O(T2)O(T^2) per head if it is materialised. FlashAttention computes it in blocks in SRAM and never materialises it → O(T)O(T) memory.

Three routes to a long context: sparse or local attention (which changes what is computed), FlashAttention (which changes how it is computed), and RAG (which changes whether it needs computing at all). The last is usually the cheapest — fetch the 5 000 relevant tokens instead of feeding in 200 000.

Code

def cost(T, d=4096, layers=32, h_kv=8, d_head=128, bytes_per_number=2):
    attn_flops = 2 * T**2 * d * layers
    lin_flops  = (4 * d**2 + 8 * d**2) * T * layers
    kv_gb = 2 * layers * h_kv * d_head * T * bytes_per_number / 1e9
    return {"T": T, "attn_TFLOPs": round(attn_flops / 1e12, 1),
            "linear_TFLOPs": round(lin_flops / 1e12, 1),
            "attn_share": round(attn_flops / (attn_flops + lin_flops), 2),
            "kv_cache_GB": round(kv_gb, 2)}

for T in (2_048, 8_192, 32_768, 131_072):
    print(cost(T))
# {'T': 2048,   'attn_TFLOPs': 1.1,   'linear_TFLOPs': 13.2, 'attn_share': 0.08, 'kv_cache_GB': 0.13}
# {'T': 8192,   'attn_TFLOPs': 17.6,  'linear_TFLOPs': 52.8, 'attn_share': 0.25, 'kv_cache_GB': 0.54}
# {'T': 32768,  'attn_TFLOPs': 281.5, 'linear_TFLOPs': 211.1,'attn_share': 0.57, 'kv_cache_GB': 2.15}
# {'T': 131072, 'attn_TFLOPs': 4504,  'linear_TFLOPs': 844,  'attn_share': 0.84, 'kv_cache_GB': 8.59}

The table shows the tipping point: at 2 k the attention is 8 % of the work, at 128 k it is 84 %.

Mastery means

  • Computes the attention cost as a function of the sequence length
  • Tells computation from memory
  • Justifies why a long context is expensive

Sign in to do the exercises and build your mastery up.

Sources

All the sources and licences