FlashAttention: Avoiding the Full Attention Matrix

黎 浩然/ 9 10 月, 2026/ 大语言模型/LARGELANGUAGEMODEL/LLM, 机器学习/MACHINELEARNING, 研究生/POSTGRADUATE, 计算机/COMPUTER/ 0 comments

Attention can create an N×N intermediate matrix whose storage grows rapidly with sequence length. FlashAttention changes data movement and intermediate storage while computing the same dense-attention definition. It does not, by default, remove token relationships.

Contents
  1. The intermediate matrix comes first
  2. Why averaging block softmaxes fails
  3. Check three scores
  4. Exact attention does not mean bitwise equality
  5. Sources

The intermediate matrix comes first

A conventional expression is O = softmax(QKᵀ / √d)V, where N is sequence length and d is head dimension. A materialized score matrix for one head holds N² elements. For N=8192 and two bytes per element, that matrix alone takes 128 MiB. This is not a total-model memory estimate.

Materialized intermediates
Full N×N scores and probabilities
Process tiles
Local scores → update statistics
Keep output
No full probability matrix stored
Original data-flow diagram showing intermediate-storage choices, not a hardware layout or bandwidth benchmark.

The original FlashAttention paper appeared in 2022. It tiles the inputs and uses on-chip storage to reduce GPU high-bandwidth-memory traffic, with selected recomputation in the backward pass. This article focuses on a forward mathematical mechanism rather than reproducing training results.

Why averaging block softmaxes fails

Normalizing each block separately and averaging its result is not global softmax: the denominators differ. Correct aggregation must retain exponential sums relative to a shared maximum.

Keep a running maximum m, an unnormalized total l and weighted sum u. For a new score a and scalar value v, let m′=max(m,a). Rescale old statistics by exp(m−m′), and give the new term weight exp(a−m′):

l′ = exp(m − m′) l + exp(a − m′)
u′ = exp(m − m′) u + exp(a − m′) v
output = u′ / l′

Subtracting the maximum keeps exponents nonpositive, avoiding direct exponentiation of large positive scores. For vector values, each component of u uses the same scale. The one-item recurrence is a minimal explanation; an actual kernel works with tiles and matrix multiplication.

Check three scores

Invented scores [0,1,2] and values [10,20,30] produce approximately 25.752104 using both a full-weight calculation and running statistics. Adding 1000 to every score preserves the result, checking softmax’s invariance to a common shift.

import math

def online(scores, values):
    m, total, weighted = -math.inf, 0.0, 0.0
    for score, value in zip(scores, values):
        new_m = max(m, score)
        old_scale = math.exp(m-new_m)
        weight = math.exp(score-new_m)
        total = old_scale*total + weight
        weighted = old_scale*weighted + weight*value
        m = new_m
    return weighted / total

def reference(scores, values):
    m = max(scores)
    weights = [math.exp(s-m) for s in scores]
    return sum(w*v for w,v in zip(weights,values))/sum(weights)

scores, values = [0.0,1.0,2.0], [10.0,20.0,30.0]
a, b = online(scores,values), reference(scores,values)
assert math.isclose(a,b,rel_tol=1e-12)
assert math.isclose(online([s+1000 for s in scores],values),b,
                    rel_tol=1e-12)
print(round(a,6))

The code was executed, with an additional 390 random cases of lengths 1 through 39 checked against the reference. It verifies one unmasked scalar weighted sum without dropout. It is not a FlashAttention kernel and checks neither gradients nor GPU or model performance.

Exact attention does not mean bitwise equality

Exact refers to retaining the dense-attention mathematical definition. A different floating-point operation order can change numerical results. Precision, masks, dropout, head dimensions and hardware support must be checked for the actual implementation.

General dense attention still includes pairwise token relationships. Less intermediate storage does not make its arithmetic linear. Compare both peak memory and latency while keeping shapes, precision and masks consistent.

This differs from PagedAttention, which the linked article discusses as KV-cache organization. Here the focus is data flow inside the attention operator. The optimizations target different layers of the serving system and are not interchangeable.

中文版

Sources

Original FlashAttention preprint (2022); authors’ implementation. Check installation and hardware support for the version actually used.

Publication note: actually published as a catch-up on October 10, 2026 (Beijing time), retaining the originally planned article date.

Support

If this article helped you, you can support this site.

WeChat support QR code; click to enlarge
WeChat
Alipay support QR code; click to enlarge
Alipay
Buy Me a Coffee; support this site
Buy Me a Coffee

Click a QR code to enlarge. More options: support page。

Share this Post

Leave a Comment

您的邮箱地址不会被公开。 必填项已用 * 标注

*
*