Back to all posts
Scaling Long-Context Attention · Part I

Exact Blocks, Rings, and Sparsity

FlashAttention, Ring Attention, and sparse attention all talk about blocks. But a block is only a unit of work: Flash changes where exact work happens, Ring changes which device performs it, and sparsity changes which terms enter the answer.

Fu-Yun Wang · 2026 · Mathematical and systems notes
Open in Colab Run FlashAttention forward and backward step by stepThe notebook checks online softmax, the saved LSE, probability reconstruction, and Triton gradients against PyTorch.

GPU Execution Model: why tiling matters

Modern GPUs are complicated machines: they have several cache levels, schedulers, memory controllers, and many kinds of execution units. That full diagram is unnecessary here. For attention, one abstraction does most of the explanatory work.

On the accelerators used to train large models, the large off-chip global memory is usually high-bandwidth memory (HBM), shared by many streaming multiprocessors (SMs). Each SM combines powerful arithmetic units with a much smaller amount of on-chip storage: a register file and shared memory implemented with SRAM. HBM is not slow in an absolute sense; its bandwidth is one reason GPUs are fast. But it is farther from the arithmetic units and more expensive to access than registers or shared SRAM. The useful picture is therefore one large, comparatively slow store connected to many compute-rich units with small, fast working spaces.

A simplified GPU diagram with one large shared HBM store above several streaming multiprocessors, each containing a small fast register and shared-SRAM working set beside high-throughput compute units.
The deliberately simplified GPU model used in this article. Real devices contain additional caches and control hardware; see NVIDIA's CUDA programming model. The HBM-versus-on-chip emphasis follows the FlashAttention cost model.

An SM therefore works on tiles. It loads a small tile from HBM, reuses that tile while it is nearby, and eventually stores a result back to HBM. The whole tensor need not fit on chip, but performance depends heavily on how often its pieces cross that boundary. Many GPU kernels are best understood as different schedules for this load–compute–store loop.

1. Dense Attention: materializing \(S\) and \(P\)

For one attention head, let \(Q,K,V\in\mathbb R^{L\times d}\). Standard attention is

\[ S=\frac{QK^\top}{\sqrt d},\qquad P_{ij}=\frac{e^{S_{ij}}}{\sum_k e^{S_{ik}}},\qquad O=PV. \] (1)

Batches and multiple heads are omitted for now: they add indices but do not change the memory argument. Both \(S\) and \(P\) have shape \(L\times L\). There are \(L^2\) query–key pairs and each dot product costs \(O(d)\), so exact dense attention performs \(\Theta(L^2d)\) arithmetic.

How conventional attention materializes \(S\) and \(P\)

No SM needs to hold the entire score matrix. The matrix multiplications are themselves tiled: many SMs load small pieces of \(Q\) and \(K\), compute score tiles, and collectively produce \(S\). In a conventional three-stage implementation, however, the complete intermediate matrices still live in HBM between kernels:

  1. the \(QK^\top\) kernel computes score tiles and writes \(S\) to HBM;
  2. the softmax kernel reads \(S\), normalizes each row, and writes \(P\);
  3. the \(PV\) kernel reads \(P\) and multiplies it by \(V\) to produce \(O\).
\[ \underbrace{\operatorname{write}(S)}_{L^2} +\underbrace{\operatorname{read}(S)}_{L^2} +\underbrace{\operatorname{write}(P)}_{L^2} +\underbrace{\operatorname{read}(P)}_{L^2}. \]

Even this idealized baseline performs at least four \(L^2\)-sized intermediate transfers across the HBM boundary. The count deliberately omits the smaller \(O(Ld)\) reads and writes of \(Q,K,V,O\); a real softmax may also take additional passes. The point is not that the GPU cannot compute a large attention matrix. It can compute it tile by tile. The inefficiency is materializing \(S\) and \(P\) in HBM and moving them again for the next stage.

FlashAttention changes the schedule, not the answer

FlashAttention uses a more careful schedule. It keeps one query tile and a small running softmax state on chip, streams through key/value tiles, computes one temporary score tile, immediately folds that tile into the running normalization and output, and then discards it. After all key/value tiles have been visited, it writes the final output tile to HBM; training also saves only \(O(L)\) row statistics for backward. The full \(S\) and \(P\) matrices are never materialized there.

This does not make dense attention subquadratic: it still evaluates every allowed query–key pair and performs the same \(\Theta(L^2d)\) arithmetic. Nor is it an approximation. Its gain comes from eliminating the avoidable \(L^2\)-sized intermediate traffic. Online softmax is the algebraic trick that makes this streaming schedule exact; Section 2 will derive it.

This is the first meaning of a block: a tile chosen so that a useful working set fits on chip. Ring Attention and sparse attention reuse the same computational unit, but give “block” different system and mathematical roles.

A three-row comparison of causal attention contracts. Diagonal blocks retain an intra-tile triangular mask. FlashAttention streams all causally required block pairs through on-chip memory; Ring Attention rotates KV blocks while skipping fully masked work; sparse attention routes a selected support and either renormalizes on it or estimates a full omitted softmax state from the query, compressed KV information, and router statistics.
Each local \(Q\times KV\) tile returns a mergeable softmax state \((m,\ell,A)\), and diagonal tiles retain an intra-tile causal mask. Flash and Ring merge all causally required block pairs into exact causal attention. Sparse routing instead changes the support: it can renormalize exactly on that selected support, or combine its exact selected state with an estimated omitted state.
The distinction to keep

FlashAttention: same pairs and same softmax, but a better memory schedule. Ring Attention: same pairs and same softmax, but KV blocks move across devices. Block-sparse attention: some pairs are not evaluated exactly, so the mathematical result changes unless their contribution is approximated elsewhere.

2. FlashAttention: exact online softmax

Take one query row. Instead of loading every logit at once, split its keys and values into blocks. Suppose earlier blocks have already been summarized by three quantities:

A new block arrives with logits \(s_k\) and values \(V_k\). Its largest logit may exceed the old maximum, so first choose the new common scale

\[ m'=\max\!\left(m,\max_{k\in\mathrm{new}} s_k\right). \]

The old sums were measured relative to \(m\), while the new block must be measured relative to \(m'\). Multiplying the old sums by \(e^{m-m'}\) simply changes units:

\[ \boxed{\begin{aligned} \ell'&=e^{m-m'}\ell+\sum_{k\in\mathrm{new}} e^{s_k-m'},\\ A'&=e^{m-m'}A+\sum_{k\in\mathrm{new}} e^{s_k-m'}V_k. \end{aligned}} \] (2)

To see why this is exact, expand the rescaled old-state contribution:

\[ e^{m-m'}\ell =e^{m-m'}\sum_{k\in\mathrm{old}}e^{s_k-m} =\sum_{k\in\mathrm{old}}e^{s_k-m'}. \]

The new-block terms in Equation (2) already use the same offset \(m'\), so the two pieces add to the exact numerator and denominator over everything seen so far.

The same reasoning can be stated as a merge invariant. After processing an index set \(\mathcal I\),

\[ m_{\mathcal I}=\max_{k\in\mathcal I}s_k,\qquad \ell_{\mathcal I} =e^{-m_{\mathcal I}}\sum_{k\in\mathcal I}e^{s_k}, \] \[ A_{\mathcal I} =e^{-m_{\mathcal I}}\sum_{k\in\mathcal I}e^{s_k}V_k. \] (2a)

For disjoint sets \(\mathcal I\) and \(\mathcal J\), choose \(m'=\max(m_{\mathcal I},m_{\mathcal J})\), rescale both states to that common maximum, and add their \(\ell\) and \(A\) terms. The result satisfies the same invariant for \(\mathcal I\cup\mathcal J\). Thus \((m,\ell,A)\) is a mergeable sufficient statistic for the exact partial numerator and denominator—the same algebra Ring Attention will reuse across devices.

A three-logit numerical check

Suppose the first tile contains logits \(1000\) and \(999\), with scalar values \(2\) and \(4\). Direct exponentiation would overflow, but the tile state is

\[ m=1000,\qquad \ell=1+e^{-1},\qquad A=2+4e^{-1}. \]

A second tile contains logit \(1001\) and value \(10\). The new maximum is \(m'=1001\), so the old state is multiplied by \(e^{-1}\):

\[ \ell'=e^{-1}(1+e^{-1})+1,\qquad A'=e^{-1}(2+4e^{-1})+10. \]

Dividing \(A'/\ell'\) is exactly the softmax-weighted average of the three values at logits \(999,1000,1001\). The large common offset never appears in an exponential.

After rescaling, the old block and new block are expressed with the same exponent offset, so their sums can be added normally. After the final block,

\[ \boxed{O=\frac{A}{\ell}},\qquad \operatorname{LSE}=m+\log\ell. \] (3)

Nothing has been dropped or approximated. The blocks are only a numerically stable way to merge pieces of the same softmax.

A LaTeX figure showing a resident query tile and online softmax state while key-value tiles stream during forward, followed by two complementary tiled traversals during backward.
Forward keeps \(Q_i,m_i,\ell_i,A_i\) resident while streaming \(K_j,V_j\). Backward reconstructs each probability tile, then uses one traversal owned by key blocks and another owned by query blocks.

What the forward kernel keeps on chip

In a minimal Triton program, one program instance owns a query tile. The query tile and the three online-softmax states live for the whole loop; each key/value tile dies after one iteration. That lifetime—not a special approximation—is the essence of the memory saving.

Annotated minimal Triton forward
import torch
import triton
import triton.language as tl

@triton.jit
def _flash_fwd(Q, K, V, O, LSE,
               stride_b: tl.constexpr, stride_h: tl.constexpr,
               H: tl.constexpr, L: tl.constexpr, D: tl.constexpr,
               SCALE: tl.constexpr, CAUSAL: tl.constexpr,
               BM: tl.constexpr, BN: tl.constexpr):
    tl.static_assert((D == 64) | (D == 128))
    q_block = tl.program_id(0)
    bh = tl.program_id(1)
    b, h = bh // H, bh % H
    base = b * stride_b + h * stride_h

    im = q_block * BM + tl.arange(0, BM)
    jn = tl.arange(0, BN)
    kd = tl.arange(0, D)
    q_ptr = Q + base + im[:, None] * D + kd[None, :]
    q = tl.load(q_ptr, mask=im[:, None] < L, other=0.0)

    # These states remain resident for the entire KV loop.
    m = tl.full([BM], -float("inf"), tl.float32)
    ell = tl.zeros([BM], tl.float32)
    acc = tl.zeros([BM, D], tl.float32)

    for start_n in range(0, L, BN):
        n = start_n + jn
        k_ptr = K + base + n[:, None] * D + kd[None, :]
        v_ptr = V + base + n[:, None] * D + kd[None, :]
        k = tl.load(k_ptr, mask=n[:, None] < L, other=0.0)
        v = tl.load(v_ptr, mask=n[:, None] < L, other=0.0)

        score = tl.dot(q, tl.trans(k)) * SCALE
        visible = (im[:, None] < L) & (n[None, :] < L)
        if CAUSAL:
            visible &= im[:, None] >= n[None, :]
        score = tl.where(visible, score, -1.0e6)

        m_new = tl.maximum(m, tl.max(score, axis=1))
        alpha = tl.exp(m - m_new)
        p = tl.exp(score - m_new[:, None])
        acc = acc * alpha[:, None] + tl.dot(p.to(q.dtype), v)
        ell = ell * alpha + tl.sum(p, axis=1)
        m = m_new

    out = acc / ell[:, None]
    o_ptr = O + base + im[:, None] * D + kd[None, :]
    tl.store(o_ptr, out, mask=im[:, None] < L)
    tl.store(LSE + bh * L + im, m + tl.log(ell), mask=im < L)

def flash_forward(q, k, v, causal=False):
    assert q.is_cuda and q.is_contiguous()
    assert q.shape == k.shape == v.shape
    assert q.dtype in (torch.float16, torch.bfloat16)
    B, H, L, D = q.shape
    assert D in (64, 128)
    o = torch.empty_like(q)
    lse = torch.empty((B, H, L), device=q.device, dtype=torch.float32)
    grid = (triton.cdiv(L, 64), B * H)
    _flash_fwd[grid](
        q, k, v, o, lse, q.stride(0), q.stride(1),
        H=H, L=L, D=D, SCALE=D ** -0.5, CAUSAL=causal,
        BM=64, BN=64, num_warps=4, num_stages=2)
    return o, lse

TileLang can express the same algorithm with explicit shared-memory buffers, fragments, copies, GEMMs, reductions, and pipeline annotations. The programming model changes; Equations (2)–(3) do not.

The same recurrence in compact TileLang form
import tilelang
import tilelang.language as T

@tilelang.jit(out_idx=[3])
def flash_tile(batch, heads, length, dim, causal,
               input_dtype="float16", BM=64, BN=64,
               stages=2, threads=128):
    assert dim in (64, 128)
    dtype = T.float16 if input_dtype == "float16" else T.bfloat16
    scale = dim ** -0.5 * 1.44269504  # exp2 scale
    shape = [batch, heads, length, dim]

    @T.prim_func
    def main(Q: T.Tensor(shape, dtype),
             K: T.Tensor(shape, dtype),
             V: T.Tensor(shape, dtype),
             O: T.Tensor(shape, dtype)):
        with T.Kernel(T.ceildiv(length, BM), heads, batch,
                      threads=threads) as (qb, h, b):
            q = T.alloc_shared([BM, dim], dtype)
            k = T.alloc_shared([BN, dim], dtype)
            v = T.alloc_shared([BN, dim], dtype)
            out = T.alloc_shared([BM, dim], dtype)
            score = T.alloc_fragment([BM, BN], T.float32)
            prob = T.alloc_fragment([BM, BN], dtype)
            acc = T.alloc_fragment([BM, dim], T.float32)
            m = T.alloc_fragment([BM], T.float32)
            m_old = T.alloc_fragment([BM], T.float32)
            alpha = T.alloc_fragment([BM], T.float32)
            ell = T.alloc_fragment([BM], T.float32)
            tile_sum = T.alloc_fragment([BM], T.float32)

            T.copy(Q[b, h, qb * BM:(qb + 1) * BM, :], q)
            T.fill(acc, 0)
            T.fill(ell, 0)
            T.fill(m, -T.infinity(T.float32))

            end = (T.min(T.ceildiv(length, BN),
                         T.ceildiv((qb + 1) * BM, BN))
                   if causal else T.ceildiv(length, BN))
            for kb in T.Pipelined(end, num_stages=stages):
                T.copy(K[b, h, kb * BN:(kb + 1) * BN, :], k)
                for i, j in T.Parallel(BM, BN):
                    qi, kj = qb * BM + i, kb * BN + j
                    valid = (qi < length) and (kj < length)
                    if causal:
                        valid = valid and (qi >= kj)
                    score[i, j] = T.if_then_else(
                        valid, 0, -T.infinity(T.float32))
                T.gemm(q, k, score, transpose_B=True,
                       policy=T.GemmWarpPolicy.FullRow)

                T.copy(m, m_old)
                T.reduce_max(score, m, dim=1, clear=True)
                for i in T.Parallel(BM):
                    m[i] = T.max(m[i], m_old[i])
                    alpha[i] = T.exp2((m_old[i] - m[i]) * scale)
                for i, j in T.Parallel(BM, BN):
                    score[i, j] = T.exp2((score[i, j] - m[i]) * scale)
                T.reduce_sum(score, tile_sum, dim=1, clear=True)
                for i in T.Parallel(BM):
                    ell[i] = ell[i] * alpha[i] + tile_sum[i]
                for i, j in T.Parallel(BM, dim):
                    acc[i, j] *= alpha[i]
                T.copy(score, prob)
                T.copy(V[b, h, kb * BN:(kb + 1) * BN, :], v)
                T.gemm(prob, v, acc, policy=T.GemmWarpPolicy.FullRow)

            for i, j in T.Parallel(BM, dim):
                acc[i, j] /= ell[i]
            T.copy(acc, out)
            T.copy(out, O[b, h, qb * BM:(qb + 1) * BM, :])

    return main

3. FlashAttention Backward: recomputing \(P\)

Here \(L^2\) means memory that grows quadratically with sequence length, not the GPU's L2 cache. Start by keeping the complete forward chain in view:

\[ S=\alpha QK^\top+M,\qquad P=\operatorname{softmax}(S),\qquad O=PV, \] \[ Q,K\in\mathbb R^{L\times d_k},\qquad V,O\in\mathbb R^{L\times d_v},\qquad \alpha=\frac{1}{\sqrt{d_k}}. \]

The mask \(M\) is zero on visible pairs and \(-\infty\) on forbidden pairs. Let \(\mathcal L\) be the final loss, and use the common deep-learning shorthand

\[ dX\equiv\frac{\partial\mathcal L}{\partial X}. \]

Thus \(dO\) is the upstream gradient—PyTorch's grad_output, not a small perturbation. Backward must turn \(dO\) into \(dQ,dK,dV\).

First pretend that forward saved the full \(P\)

For query \(i\), the output is \(O_i=\sum_jP_{ij}V_j\). A value vector \(V_j\) contributes to every query output, so its gradient collects all of those paths:

\[ dV_j=\sum_iP_{ij}dO_i, \qquad\boxed{dV=P^\top dO}. \]

The scalar \(P_{ij}\) multiplies vector \(V_j\) inside \(O_i\), so its gradient is their inner product with the upstream direction:

\[ dP_{ij}=\langle dO_i,V_j\rangle=dO_iV_j^\top, \qquad\boxed{dP=dOV^\top}. \]

If forward had stored \(P\), these two steps would be immediate. But \(P\) has shape \(L\times L\). FlashAttention deliberately does not write that full matrix to HBM, so backward must recover only the probability tiles it currently needs.

Deriving the row-softmax gradient

Softmax is row-wise. Fix query row \(i\), write \(p_j=P_{ij}\) and \(s_j=S_{ij}\), and temporarily suppress the row index:

\[ p_j=\frac{e^{s_j}}{\sum_\ell e^{s_\ell}},\qquad \frac{\partial p_j}{\partial s_k} =p_j\bigl(\mathbf 1[j=k]-p_k\bigr). \]

Every \(p_j\) depends on \(s_k\) through the shared denominator, so every path contributes to \(ds_k\):

\[ \begin{aligned} ds_k &=\sum_jdp_j\frac{\partial p_j}{\partial s_k}\\ &=\sum_jdp_jp_j\bigl(\mathbf 1[j=k]-p_k\bigr)\\ &=dp_kp_k-p_k\sum_jp_jdp_j\\ &=p_k\left(dp_k-\sum_jp_jdp_j\right). \end{aligned} \]

Restoring the row and column indices, define one scalar per query row:

\[ D_i=\sum_jP_{ij}dP_{ij},\qquad \boxed{dS_{ij}=P_{ij}\bigl(dP_{ij}-D_i\bigr)}. \] (4)

The subtraction handles the coupling introduced by normalization. For example, if every \(dP_{ij}\) in a row equals the same constant \(c\), then \(D_i=\sum_jP_{ij}c=c\), so every \(dS_{ij}=P_{ij}(c-c)=0\). This is exactly right: because \(\sum_jP_{ij}=1\), assigning the same gradient to every probability creates no effective logit direction.

Why \(D_i\) does not require the saved probability row

The definition of \(D_i\) still appears to require \(P_i\). Substitute the value-mixture gradient \(dP_{ij}=dO_iV_j^\top\):

\[ \begin{aligned} D_i &=\sum_jP_{ij}dP_{ij}\\ &=\sum_jP_{ij}\bigl(dO_iV_j^\top\bigr)\\ &=dO_i\left(\sum_jP_{ij}V_j\right)^\top\\ &=\boxed{dO_iO_i^\top} =\boxed{\sum_rO_{ir}dO_{ir}}. \end{aligned} \] (4a)

This removes the first apparent need for \(P\): a small preprocessing kernel computes \(D_i\) from the saved output \(O_i\) and upstream gradient \(dO_i\).

Reconstructing a probability tile without saving \(P\)

Computing \(dS_{ij}\) still requires the local probability \(P_{ij}\). Backward can cheaply recompute a score tile from \(Q\) and \(K\), but softmax normalization depends on all visible keys in that query row—not just the current tile. That missing row normalizer is exactly why forward deliberately saves LSE.

For the keys \(\mathcal A(i)\) visible to query \(i\),

\[ \operatorname{LSE}_i =\log\sum_{j\in\mathcal A(i)}e^{S_{ij}} =m_i+\log\ell_i. \]

The second equality is Equation (3): online softmax has already computed the stable row maximum \(m_i\) and rescaled sum \(\ell_i\). Saving their combined log-normalizer costs only one scalar per query row. Backward can then recover any unmasked probability directly:

\[ \begin{aligned} P_{ij} &=\frac{e^{S_{ij}}}{\sum_{\ell\in\mathcal A(i)}e^{S_{i\ell}}}\\ &=\exp\!\left(S_{ij}-\log\sum_{\ell\in\mathcal A(i)}e^{S_{i\ell}}\right)\\ &=\boxed{\exp\!\left(S_{ij}-\operatorname{LSE}_i\right)}. \end{aligned} \]

Masked positions remain zero. For query rows \(I\) and key rows \(J\), the kernel recomputes only

\[ S_{I,J}=\alpha Q_IK_J^\top+M_{I,J},\qquad P_{I,J}=\exp\!\left(S_{I,J}-\operatorname{LSE}_I[:,\mathrm{None}]\right). \]
The full \(S\) and \(P\) are never stored. Forward creates one temporary score tile, merges it into \((m,\ell,A)\), and discards it. Backward recreates one \(S\) tile and one \(P\) tile, uses them immediately, and discards them. Beyond the required inputs \(Q,K,V\), forward retains \(O\in\mathbb R^{L\times d_v}\) and the length-\(L\) LSE vector for backward; no \(L\times L\) tensor is written to HBM.

From \(dS\) to \(dQ\) and \(dK\)

Once a probability tile has been reconstructed, backward computes \(dP_{I,J}=dO_IV_J^\top\), broadcasts the precomputed \(D_I\) across its key columns, and forms \(dS_{I,J}=P_{I,J}\odot(dP_{I,J}-D_I[:,\mathrm{None}])\). Since \(M\) does not depend on \(Q,K\), the remaining matrix-product gradients are

\[ \boxed{dQ=\alpha\,dSK},\qquad \boxed{dK=\alpha\,dS^\top Q}. \] \[ dQ_i=\alpha\sum_jdS_{ij}K_j,\qquad dK_j=\alpha\sum_idS_{ij}Q_i. \] (5)

Putting the backward tile computation together

The preprocessing step computes \(D_i=dO_iO_i^\top\) once per query row. Each subsequent query–KV tile interaction then follows the same short chain:

  1. reload the current \(Q,K,V,dO\) tiles;
  2. recompute \(S=\alpha QK^\top+M\);
  3. load the saved row LSE and reconstruct \(P=\exp(S-\operatorname{LSE})\);
  4. compute \(dP=dOV^\top\) and \(dS=P\odot(dP-D)\);
  5. accumulate this tile's contributions to \(dQ,dK,dV\), then discard \(S,P,dP,dS\).

The extra \(QK^\top\) and exponential work replaces writing and rereading quadratic intermediates. On modern GPUs, that trade is usually favorable because matrix multiplication is fast while moving an \(L\times L\) tensor through HBM is expensive.

Backward must use the same scale and mask as forward. Otherwise LSE normalizes a different set of logits and the kernel computes the gradient of a different function.

The ownership conflict between \(dQ\) and \(dK,dV\)

Gradient tileReduction directionRace-free owner
\(dQ_I\)All KV tiles \(J\)One program owns query tile \(I\)
\(dK_J\)All query tiles \(I\)One program owns key tile \(J\)
\(dV_J\)All query tiles \(I\)One program owns value tile \(J\)

The official Triton fused-attention tutorial resolves the opposing reductions with two logical traversals inside its backward kernel: a KV-owned traversal finishes \(dK,dV\), and a query-owned traversal finishes \(dQ\). This recomputes local \(S,P,dP,dS\) twice but avoids atomics. It is an implementation choice, not a law of FlashAttention: FlashAttention-2 instead uses KV ownership and atomically accumulates partial \(dQ\).

Runnable Triton backward: preprocess, dK/dV, then dQ
@triton.jit
def _flash_bwd_preprocess(O, DO, Delta,
                          stride_b: tl.constexpr, stride_h: tl.constexpr,
                          H: tl.constexpr, L: tl.constexpr, D: tl.constexpr,
                          BLOCK: tl.constexpr):
    block = tl.program_id(0)
    bh = tl.program_id(1)
    b, h = bh // H, bh % H
    base = b * stride_b + h * stride_h
    rows = block * BLOCK + tl.arange(0, BLOCK)
    dims = tl.arange(0, D)

    o = tl.load(O + base + rows[:, None] * D + dims[None, :],
                mask=rows[:, None] < L, other=0.0)
    do = tl.load(DO + base + rows[:, None] * D + dims[None, :],
                 mask=rows[:, None] < L, other=0.0).to(tl.float32)
    delta = tl.sum(o * do, axis=1)  # D_i = <O_i, dO_i>
    tl.store(Delta + bh * L + rows, delta, mask=rows < L)


@triton.jit
def _flash_bwd(Q, K, V, DO, DQ, DK, DV, LSE, Delta,
               stride_b: tl.constexpr, stride_h: tl.constexpr,
               H: tl.constexpr, L: tl.constexpr, D: tl.constexpr,
               SCALE: tl.constexpr, CAUSAL: tl.constexpr,
               BLOCK: tl.constexpr):
    tile = tl.program_id(0)
    bh = tl.program_id(1)
    b, h = bh // H, bh % H
    base = b * stride_b + h * stride_h
    in_tile = tl.arange(0, BLOCK)
    dims = tl.arange(0, D)

    # Traversal 1: this program owns one KV tile.
    cols = tile * BLOCK + in_tile
    k = tl.load(K + base + cols[:, None] * D + dims[None, :],
                mask=cols[:, None] < L, other=0.0)
    v = tl.load(V + base + cols[:, None] * D + dims[None, :],
                mask=cols[:, None] < L, other=0.0)
    dk = tl.zeros([BLOCK, D], tl.float32)
    dv = tl.zeros([BLOCK, D], tl.float32)

    for start_m in range(0, L, BLOCK):
        rows = start_m + in_tile
        q = tl.load(Q + base + rows[:, None] * D + dims[None, :],
                    mask=rows[:, None] < L, other=0.0)
        do = tl.load(DO + base + rows[:, None] * D + dims[None, :],
                     mask=rows[:, None] < L, other=0.0)
        lse = tl.load(LSE + bh * L + rows, mask=rows < L, other=0.0)
        delta = tl.load(Delta + bh * L + rows, mask=rows < L, other=0.0)

        scores = tl.dot(q, tl.trans(k)) * SCALE
        visible = (rows[:, None] < L) & (cols[None, :] < L)
        if CAUSAL:
            visible = visible & (rows[:, None] >= cols[None, :])
        scores = tl.where(visible, scores, -1.0e6)
        p = tl.where(visible, tl.exp(scores - lse[:, None]), 0.0)
        dp = tl.dot(do, tl.trans(v)).to(tl.float32)
        ds = p * (dp - delta[:, None])

        dv += tl.dot(tl.trans(p.to(q.dtype)), do)
        dk += tl.dot(tl.trans(ds.to(q.dtype)), q) * SCALE

    tl.store(DK + base + cols[:, None] * D + dims[None, :], dk,
             mask=cols[:, None] < L)
    tl.store(DV + base + cols[:, None] * D + dims[None, :], dv,
             mask=cols[:, None] < L)

    # Traversal 2: the same program id now owns one query tile.
    rows = tile * BLOCK + in_tile
    q = tl.load(Q + base + rows[:, None] * D + dims[None, :],
                mask=rows[:, None] < L, other=0.0)
    do = tl.load(DO + base + rows[:, None] * D + dims[None, :],
                 mask=rows[:, None] < L, other=0.0)
    lse = tl.load(LSE + bh * L + rows, mask=rows < L, other=0.0)
    delta = tl.load(Delta + bh * L + rows, mask=rows < L, other=0.0)
    dq = tl.zeros([BLOCK, D], tl.float32)

    for start_n in range(0, L, BLOCK):
        cols = start_n + in_tile
        k = tl.load(K + base + cols[:, None] * D + dims[None, :],
                    mask=cols[:, None] < L, other=0.0)
        v = tl.load(V + base + cols[:, None] * D + dims[None, :],
                    mask=cols[:, None] < L, other=0.0)

        scores = tl.dot(q, tl.trans(k)) * SCALE
        visible = (rows[:, None] < L) & (cols[None, :] < L)
        if CAUSAL:
            visible = visible & (rows[:, None] >= cols[None, :])
        scores = tl.where(visible, scores, -1.0e6)
        p = tl.where(visible, tl.exp(scores - lse[:, None]), 0.0)
        dp = tl.dot(do, tl.trans(v)).to(tl.float32)
        ds = p * (dp - delta[:, None])
        dq += tl.dot(ds.to(q.dtype), k) * SCALE

    tl.store(DQ + base + rows[:, None] * D + dims[None, :], dq,
             mask=rows[:, None] < L)


def flash_backward(q, k, v, o, lse, do, causal=False):
    B, H, L, D = q.shape
    dq, dk, dv = torch.empty_like(q), torch.empty_like(k), torch.empty_like(v)
    delta = torch.empty((B, H, L), device=q.device, dtype=torch.float32)
    block = 64
    grid = (triton.cdiv(L, block), B * H)

    _flash_bwd_preprocess[grid](
        o, do, delta, q.stride(0), q.stride(1),
        H=H, L=L, D=D, BLOCK=block, num_warps=4)
    _flash_bwd[grid](
        q, k, v, do, dq, dk, dv, lse, delta,
        q.stride(0), q.stride(1),
        H=H, L=L, D=D, SCALE=D ** -0.5, CAUSAL=causal,
        BLOCK=block, num_warps=4, num_stages=2)
    return dq, dk, dv

FlashAttention has now solved the quadratic-intermediate problem on one accelerator: score and probability tiles are streamed through a small on-chip working set instead of being materialized in HBM. It does not eliminate memory that grows linearly with the sequence, however. The \(Q,K,V\) tensors and training activations still grow with \(L\), so a sufficiently long sequence eventually exceeds one device's memory. The mergeable softmax state suggests the next step: keep a query block local, but stream KV blocks across devices rather than only from local HBM.

4. Ring Attention: rotating KV blocks across devices

This limit is becoming practical as training horizons grow far beyond ordinary document lengths. A long-running agent can accumulate tool traces, observations, actions, and memory into a trajectory approaching millions of tokens. Video expands even faster under patch tokenization. For illustration, 30 seconds of 1280-by-720 video at 24 frames per second produces

\[ \frac{1280}{16}\times\frac{720}{16}\times \frac{30\times24}{2} =80\times45\times360 \approx1.30\text{ million tokens} \]

under 16-by-16 spatial patches and two-frame tubelets. This is an order-of-magnitude example, not a universal conversion: video VAEs, spatial and temporal downsampling, patch sizes, and tubelet sizes all change the actual token count. The systems problem remains the same once the resulting sequence no longer fits comfortably on one accelerator.

Ring Attention partitions that sequence across \(G\) devices. Device \(g\) keeps its local query block \(Q_g\) stationary. Its current \(K,V\) block is the moving object: the device computes attention between \(Q_g\) and that block, merges the result into its online-softmax state, and passes the KV block one hop to the next device. The blocks move in one direction around the ring rather than back and forth. After \(G\) steps, every query block has met every allowed KV block.

\[ (m_g,\ell_g,A_g) \xrightarrow{(K_0,V_0)} \xrightarrow{(K_1,V_1)}\cdots \xrightarrow{(K_{G-1},V_{G-1})} O_g=\frac{A_g}{\ell_g}. \]

Each step uses the same merge rule in Equation (2). Communication changes the order in which KV blocks arrive, but exact partial softmax statistics can be merged in any block order. Ring Attention therefore extends the sequence across devices and overlaps communication with block computation; it does not reduce the total dense pair set by itself.

More formally, suppose two devices or two ring steps produce partial states \((m_a,\ell_a,A_a)\) and \((m_b,\ell_b,A_b)\) for disjoint KV blocks. Their associative merge is

\[ \begin{aligned} m_{a\cup b}&=\max(m_a,m_b),\\ \ell_{a\cup b} &=e^{m_a-m_{a\cup b}}\ell_a +e^{m_b-m_{a\cup b}}\ell_b,\\ A_{a\cup b} &=e^{m_a-m_{a\cup b}}A_a +e^{m_b-m_{a\cup b}}A_b. \end{aligned} \] (5a)

Because this operator summarizes the union of two index sets, merging blocks in ring order or local-memory order produces the same real-arithmetic result. A causal implementation can skip blocks that lie entirely in the future and mask the one block that crosses the diagonal. Communication is useful only when sending the next KV block is overlapped by enough attention work on the current block; otherwise the exact mathematics remains correct but scaling stalls on the interconnect.

“Block attention” is not automatically sparse attention FlashAttention blocks an exact matrix for memory locality. Ring Attention distributes those exact blocks. A block-sparse method is different because some allowed blocks never contribute exactly to the numerator or denominator.

5. Sparse Attention: selecting fewer query–key pairs

FlashAttention gives us something useful beyond a faster dense kernel: it packages attention into \(Q\)-tile by \(KV\)-tile units. Dense FlashAttention still visits every allowed tile, but if an entire KV tile is unlikely to matter for a query tile, the schedule can omit it and avoid its memory loads and matrix multiplications altogether.

Why might this help on a long sequence? Attention is often concentrated rather than uniform. After softmax, a small number of tokens or blocks may carry most of the probability mass, while many distant tokens contribute almost nothing. As the context grows, computing every query–key pair becomes increasingly wasteful if only a small fraction materially changes the output.

There is an immediate catch: the important positions are unknown until the query is compared with the keys. Computing the full \(QK^\top\) matrix merely to decide what to skip would save nothing. Sparse attention therefore needs a selector (or router): a cheaper scoring mechanism that predicts which tokens or blocks deserve the full attention computation.

The selected items are then evaluated exactly with their full keys and values. For everything else, a method must make a second decision: drop it, preserve a coarse version through another path, or approximate its missing contribution. These are separate questions—where should exact computation go? and what should happen to what was not selected?

The approximation can now be stated precisely. For one query \(i\), let \(\mathcal A_i\) be all keys allowed by causal and semantic masks, and define \(s_{ij}=q_i^\top k_j/\sqrt d\). Dense attention is one fraction:

\[ D_i=\sum_{j\in\mathcal A_i}e^{s_{ij}},\qquad N_i=\sum_{j\in\mathcal A_i}e^{s_{ij}}v_j,\qquad o_i=\frac{N_i}{D_i}. \] (6)

Partition the allowed keys into blocks \(B_1,\ldots,B_M\). This alone introduces no approximation:

\[ D_i=\sum_bD_{ib},\quad D_{ib}=\sum_{j\in B_b}e^{s_{ij}}, \qquad N_i=\sum_bN_{ib},\quad N_{ib}=\sum_{j\in B_b}e^{s_{ij}}v_j. \]

The selector chooses an exact set \(E_i\), often a union of selected blocks and a mandatory local window. The omitted set is \(U_i=\mathcal A_i\setminus E_i\), so

\[ N_i=N_{i,E}+N_{i,U},\qquad D_i=D_{i,E}+D_{i,U}. \] (7)

This split reveals two possible errors. The selection error comes from failing to place important tokens in \(E_i\). The aggregation error comes from how the method handles \(N_{i,U}\) and \(D_{i,U}\) after selection. A good router does not answer the second question.

Why the sparse unit is usually a block

FlashAttention has already turned dense attention into a loop over query tiles and KV tiles. Let \(Q_a\in\mathbb R^{B_Q\times d}\) be one query tile and \((K_b,V_b)\) one KV tile. A dense kernel visits every causally allowed tile \(b\), computes

\[ S_{ab}=\frac{Q_aK_b^\top}{\sqrt d},\qquad m_{ab}=\operatorname{rowmax}(S_{ab}), \] \[ \ell_{ab}=\operatorname{rowsum}e^{S_{ab}-m_{ab}},\qquad A_{ab}=e^{S_{ab}-m_{ab}}V_b, \] \[ (m_a,\ell_a,A_a) \leftarrow \operatorname{Merge}\!\left( (m_a,\ell_a,A_a),(m_{ab},\ell_{ab},A_{ab}) \right), \]

and merges the tile into the running online-softmax state \((m_a,\ell_a,A_a)\). Block-sparse attention changes the outer loop:

\[ \underbrace{b\in\{1,\ldots,M\}}_{\text{dense FlashAttention}} \qquad\longrightarrow\qquad \underbrace{b\in\mathcal C_a}_{\text{block-sparse attention}}, \] (7a)

If \(b\notin\mathcal C_a\), the kernel can skip loading \(K_b,V_b\), skip the \(Q_aK_b^\top\) GEMM, and skip the \(P_{ab}V_b\) GEMM. This is why deleting a whole tile can produce real speed, whereas scattering zeros inside every tile usually does not: the dense tile still has to be loaded and multiplied.

With \(n_Q=L/B_Q\) query tiles, \(n_K=L/B_K\) KV tiles, and \(r\) selected KV tiles per query tile, the number of tile pairs falls from approximately \(n_Qn_K\) to \(n_Qr\). The selector is useful only if finding \(\mathcal C_a\) and gathering its KV blocks costs less than the skipped tile work.

Inside every selected tile, the computation is still ordinary FlashAttention. The sparse approximation lies in the tile schedule: under drop-and-renormalize, only selected tiles are merged into the online-softmax numerator and denominator. NSA, HiLS, and PISA keep the same blockwise execution unit but define additional paths for coarse information, hierarchical mass, or the omitted tail.

Not every selector is block-granular DSA assigns a score to every previous token and selects individual token positions with top-\(k\); it is logically token-wise even if its kernel batches the resulting work. NSA, HiLS, PISA, and many GPU-oriented sparse kernels instead select contiguous blocks because those blocks map directly to FlashAttention-style loads and GEMMs. “Sparse” describes the chosen pair set; “block sparse” additionally requires that set to be selected in block-shaped units.

The selector must score something cheaper than the full \(QK^\top\) matrix. It may use a fixed block statistic, a learned compressed block, or a separate low-dimensional token indexer.

Fixed pooling: a cheap proxy, not the true block mass

For a mean-pooled block proxy, define

\[ \bar k_b=\frac{1}{|B_b|}\sum_{j\in B_b}k_j,\qquad r_{ib}=\frac{q_i^\top\bar k_b}{\sqrt d},\qquad \mathcal C_i=\operatorname{TopK}_b(r_{ib},m). \] (7b)

Let \(B=|B_b|\), and write the mean logit in block \(b\) as

\[ \bar s_{ib} =\frac1B\sum_{j\in B_b}s_{ij} =\frac{q_i^\top\bar k_b}{\sqrt d}. \]

What is the block’s “true score”? For one query, \(s_{ij}\) is token \(j\)’s logit before softmax, so \(e^{s_{ij}}\) is its unnormalized attention weight. The whole block contributes

\[ A_{ib}=\sum_{j\in B_b}e^{s_{ij}}, \qquad f_b(q_i)=\log A_{ib}. \]

To see when mean pooling is accurate, write each logit as its block mean plus a deviation:

\[ s_{ij}=\bar s_{ib}+\epsilon_{ij}, \qquad \frac1B\sum_{j\in B_b}\epsilon_{ij}=0. \] \[ \boxed{ f_b(q_i) =\bar s_{ib}+\log B+ \log\!\left(\frac1B \sum_{j\in B_b}e^{\epsilon_{ij}}\right) }. \]

If the logits inside the block are similar, every \(\epsilon_{ij}\) is close to zero. Then every \(e^{\epsilon_{ij}}\) is close to \(1\), so the last logarithm is close to zero. For equal-sized blocks, \(\log B\) is also the same constant. The exact block score is therefore approximately \(\bar s_{ib}+\log B\), and ranking by the pooled score \(\bar s_{ib}\) gives almost the same order.

The approximation becomes worse as the logits spread out. For small deviations, the omitted term is approximately

\[ \log\!\left(\frac1B\sum_j e^{\epsilon_{ij}}\right) \approx\frac{1}{2B}\sum_j\epsilon_{ij}^2. \]

So mean pooling is reliable for a homogeneous block, but it can miss a block containing one unusually relevant token: exponentiation gives that token much more weight, while the arithmetic mean dilutes it among the other tokens.

DSA: a distilled per-token indexer

DSA treats selection as a teacher–student problem. Dense attention is the teacher; a lighter, token-wise lightning indexer is the student. Instead of using the expensive main-attention scores to choose tokens, the indexer projects each token into a cheaper indexing space and predicts which positions the teacher would consider important.

Each previous token \(j\) has one index key \(k^I_j\in\mathbb R^{d_I}\). Query \(i\) has \(H^I\) indexer query heads \(q^I_{i,a}\in\mathbb R^{d_I}\) and scalar weights \(w^I_{i,a}\). Their votes are combined into one score for token \(j\):

\[ I_{ij} =\sum_{a=1}^{H^I}w^I_{i,a} \operatorname{ReLU}\!\left((q^I_{i,a})^\top k^I_j\right), \qquad E_i=\operatorname{TopK}_j(I_{ij},k). \] (7c)

The indexer is therefore an attention-like scorer with fewer heads and a cheaper feature space than DSA's main attention computation. It is not a smaller attention layer. It has no value aggregation and produces no contextualized output; all indexer heads only vote on one token-ranking score, after which top-\(k\) returns the positions that the main attention should evaluate exactly.

The top-\(k\) operation is a discrete choice: small score changes usually leave the selected set unchanged, and crossing a selection boundary changes it abruptly. Ordinary backpropagation through membership therefore does not provide a useful training signal for the indexer. DSA supplies an explicit target instead.

During dense warm-up, the model still runs full attention. Let \(p^{(h)}_{ij}\) be its attention probability from head \(h\). DSA sums the probabilities across heads and normalizes them along the sequence to obtain one teacher distribution \(t_i\). The indexer is trained to match that distribution:

\[ t_{ij} =\frac{\sum_h p^{(h)}_{ij}} {\sum_{\ell\in\mathcal A_i}\sum_h p^{(h)}_{i\ell}}, \qquad \widehat p_i=\operatorname{Softmax}(I_i), \] \[ \boxed{\mathcal L_{\mathrm{index}} =\sum_i\operatorname{KL}(t_i\,\|\,\widehat p_i).} \] (7d)

The teacher target is therefore not one particular attention head: it represents the total importance assigned to each token across all heads. After warm-up, sparse training activates top-\(k\) retrieval and continues adapting the model and indexer. DSA then drops \(U_i\) and applies ordinary softmax only to the selected tokens.

6. Sparse Attention Semantics: accounting for unselected mass

Once \(E_i\) is fixed, several mathematically different outputs remain possible.

Drop and renormalize

\(\widetilde o_i=N_{i,E}/D_{i,E}\)

Unselected mass becomes zero, and all remaining probability is redistributed over selected tokens. This is exact softmax on \(E_i\), not the selected portion of dense softmax.

Separately normalized branches

\(\widetilde o_i=\sum_b g_{i,b}o_{i,b}\)

Compressed, selected, and local branches normalize separately, then learned gates combine their outputs. NSA uses this contract.

Hierarchical mass

\(\widetilde o_i=\sum_b\pi_{ib}\sum_j\rho_{ij\mid b}v_j\)

A block-level distribution controls mass across selected chunks, while exact logits normalize within each chunk. Unselected chunks still have zero mass.

Shared tail correction

\(\widetilde o_i=(N_E+\widehat N_U)/(D_E+\widehat D_U)\)

Selected terms are exact and omitted terms are approximated inside the same numerator and denominator. PISA explicitly follows this route.

The first contract is common in masking systems such as FlexAttention and learned selection such as DSA. Native Sparse Attention uses independently normalized compression, selection, and local branches. HiLS introduces hierarchical normalization, while PISA estimates the missing tail so exact and approximate terms share one denominator.

Exactly what drop-and-renormalize changes

Define the exact softmax output within each subset:

\[ o_{i,E}=\frac{N_{i,E}}{D_{i,E}},\qquad o_{i,U}=\frac{N_{i,U}}{D_{i,U}},\qquad \lambda_i=\frac{D_{i,E}}{D_{i,E}+D_{i,U}}. \]

The dense output can then be decomposed without approximation:

\[ \boxed{o_i=\lambda_i o_{i,E}+(1-\lambda_i)o_{i,U}.} \] (8)

Drop-and-renormalize returns \(\widetilde o_i=o_{i,E}\). Therefore its exact error relative to dense attention is

\[ \boxed{\widetilde o_i-o_i =(1-\lambda_i)(o_{i,E}-o_{i,U}).} \] (9)

Two things must be small: the omitted probability mass \(1-\lambda_i\), or the difference between selected and omitted value averages. Selecting tokens with large logits helps the first condition, but the output error also depends on values.

A tiny example makes renormalization visible. Suppose three logits are all zero and their scalar values are \(0,3,9\). Dense attention outputs \(4\). If the router selects the first two tokens, their dense probabilities were \(1/3\) each, but selected-set softmax changes them to \(1/2\) each and outputs \(1.5\). The missing token is not merely removed; its \(1/3\) probability mass is redistributed.

NSA: a compressed path for routing and memory

NSA first maps a contiguous, possibly overlapping KV block to a learned compressed token:

\[ \widetilde k_b=\phi_K(K_{B_b}),\qquad \widetilde v_b=\phi_V(V_{B_b}), \] \[ p^{\mathrm{cmp}}_{ib} =\operatorname{Softmax}_b\!\left( \frac{q_i^\top\widetilde k_b}{\sqrt d}\right), \qquad \mathcal C_i=\operatorname{TopKBlocks}(p^{\mathrm{cmp}}_i). \] (10)

The compressed score is reused to choose contiguous blocks of the original KV sequence for fine-grained attention. When compression blocks overlap or use a different size from selection blocks, their scores are aggregated by temporal overlap. Under GQA, query heads sharing one physical KV cache also share the block choice, avoiding a union of unrelated gathers.

NSA computes separately normalized compressed, selected, and sliding-window outputs, then mixes them:

\[ \widetilde o_i^{\mathrm{NSA}} =\sigma(g_{i,c})o_{i,\mathrm{cmp}} +\sigma(g_{i,s})o_{i,\mathrm{selected}} +\sigma(g_{i,w})o_{i,\mathrm{window}}. \] (10a)

The gates do not have to sum to one, and each \(o\) was produced by its own denominator. A compressed block value is therefore a real information path, not merely a routing summary. Because that path affects the language-model output, \(\phi_K,\phi_V\), and the gates can learn from the LM loss without a dense-attention teacher. Hard block membership itself is still discrete; the differentiable compressed branch is what prevents the learned proxy from being an unsupervised side calculation.

NSA therefore does not reconstruct the missing dense numerator and denominator. Unselected raw values are absent from the selected branch, while a learned coarse representation of the full history survives through \(o_{i,\mathrm{cmp}}\).

HiLS: estimating chunk mass with a landmark

HiLS inserts a landmark query \(q'_b\) for each chunk. Within that chunk, it forms

\[ p_{bj} =\frac{e^{(q'_b)^\top k_j/\sqrt d}} {\sum_{\ell\in B_b}e^{(q'_b)^\top k_\ell/\sqrt d}}, \qquad k'_b=\sum_{j\in B_b}p_{bj}k_j. \]

Recall the exact log-mass \(f_b(q)=\log\sum_{j\in B_b}e^{q^\top k_j/\sqrt d}\). Its gradient at the landmark is

\[ \left.\nabla_qf_b(q)\right|_{q=q'_b} =\frac{1}{\sqrt d}\sum_{j\in B_b}p_{bj}k_j =\frac{k'_b}{\sqrt d}. \]

A first-order Taylor expansion around \(q'_b\) gives

\[ f_b(q_i) \approx f_b(q'_b) +\frac{(q_i-q'_b)^\top k'_b}{\sqrt d} =\frac{q_i^\top k'_b}{\sqrt d} +\left[f_b(q'_b)-\frac{(q'_b)^\top k'_b}{\sqrt d}\right]. \]

The bracketed intercept looks opaque until the definition of \(p_{bj}\) is substituted:

\[ \begin{aligned} f_b(q'_b)-\frac{(q'_b)^\top k'_b}{\sqrt d} &= \sum_jp_{bj}\left[ \log\sum_\ell e^{(q'_b)^\top k_\ell/\sqrt d} -\frac{(q'_b)^\top k_j}{\sqrt d}\right]\\ &=-\sum_jp_{bj}\log p_{bj} =H(p_b). \end{aligned} \]

Therefore the router score

\[ \boxed{\widehat s_{ib} =\frac{q_i^\top k'_b}{\sqrt d}+H(p_b) \approx \log\sum_{j\in B_b}e^{s_{ij}}.} \] (11)

has a precise interpretation: \(k'_b\) is the local slope of chunk log-mass and entropy is its Taylor intercept. The remaining step is not an ordinary softmax over all retrieved tokens. HiLS factorizes the probability into an exact conditional distribution inside each selected chunk and a learned distribution across chunks:

\[ Z_{ib}=\sum_{j\in B_b}e^{s_{ij}},\qquad \rho_{ij\mid b}=\frac{e^{s_{ij}}}{Z_{ib}},\qquad \widehat Z_{ib}=e^{\widehat s_{ib}}, \] \[ \pi_{ib} =\frac{\widehat Z_{ib}} {\sum_{c\in\mathcal C_i}\widehat Z_{ic}+Z_{i,\mathrm{local}}}, \qquad \pi_{i,\mathrm{local}} =\frac{Z_{i,\mathrm{local}}} {\sum_{c\in\mathcal C_i}\widehat Z_{ic}+Z_{i,\mathrm{local}}}, \] \[ \boxed{\widetilde o_i^{\mathrm{HiLS}} =\sum_{b\in\mathcal C_i}\pi_{ib} \sum_{j\in B_b}\rho_{ij\mid b}v_j +\pi_{i,\mathrm{local}}o_{i,\mathrm{local}}.} \] (11a)

The local window contributes its exact mass \(Z_{i,\mathrm{local}}\); a retrieved distant chunk contributes its surrogate mass \(\widehat Z_{ib}\). Because selected router scores appear in \(\pi_{ib}\), the language-model loss trains the landmark representation directly. The top-\(K\) membership remains hard, however, so unselected distant chunks receive neither output mass nor this gradient. HiLS improves how probability is allocated inside the retained set; it does not restore the dense tail.

PISA: approximating the omitted tail in the same softmax

PISA keeps selected blocks exact and approximates unselected blocks inside the same numerator and denominator. Write \(k_j=\mu_b+\delta_j\), where \(\mu_b\) is a block mean. A first-order expansion gives

\[ e^{q_i^\top k_j/\sqrt d} =e^{q_i^\top\mu_b/\sqrt d} e^{q_i^\top\delta_j/\sqrt d} \approx e^{q_i^\top\mu_b/\sqrt d} \left(1+\frac{q_i^\top\delta_j}{\sqrt d}\right). \]

Summing over a block yields

\[ D_{ib}\approx e^{q_i^\top\mu_b/\sqrt d}|B_b|, \] \[ N_{ib}\approx e^{q_i^\top\mu_b/\sqrt d} \left[ \sum_{j\in B_b}v_j +\frac{q_i^\top}{\sqrt d} \sum_{j\in B_b}\delta_jv_j^\top \right]. \] (12)

Because \(\sum_j\delta_j=0\), the first-order denominator correction vanishes, but the value-weighted numerator correction generally does not. Define

\[ \alpha_{ib}=e^{q_i^\top\mu_b/\sqrt d},\qquad V_b^\Sigma=\sum_{j\in B_b}v_j,\qquad H_b=\sum_{j\in B_b}\delta_jv_j^\top. \]

Then one unselected block contributes approximately \(\alpha_{ib}V_b^\Sigma+\alpha_{ib}q_i^\top H_b/\sqrt d\) to the numerator. This expression is mathematically cheap in pair count but hardware-unfriendly if every query must stream a different \(d\times d_v\) matrix \(H_b\) for every tail block.

PISA replaces those many first-order matrices by one precomputed global statistic,

\[ \overline H=\frac1M\sum_{b=1}^{M}H_b, \] \[ \sum_{b\in U_i}\alpha_{ib} \frac{q_i^\top H_b}{\sqrt d} \;\approx\; \left(\sum_{b\in U_i}\alpha_{ib}\right) \frac{q_i^\top\overline H}{\sqrt d}. \] (12a)

The zeroth-order term still scans compact block means and value sums, but the first-order correction now loads only one shared matrix:

\[ \widehat D_{i,U}=|B|\sum_{b\in U_i}\alpha_{ib}, \] \[ \widehat N_{i,U} =\sum_{b\in U_i}\alpha_{ib}V_b^\Sigma +\left(\sum_{b\in U_i}\alpha_{ib}\right) \frac{q_i^\top\overline H}{\sqrt d}, \] \[ \boxed{\widetilde o_i^{\mathrm{PISA}} =\frac{N_{i,E}+\widehat N_{i,U}} {D_{i,E}+\widehat D_{i,U}}.} \] (12b)

Selected exact terms and approximate tail terms therefore compete in one normalization. This is mathematically different from NSA’s separately normalized compressed branch, and it is what lets PISA reuse pretrained diffusion-Transformer weights without training a new router.

The selector also accounts for where the global approximation is least trustworthy. A simplified covariance-aware block score is

\[ r_{ib}^{\mathrm{PISA}} =\frac{q_i^\top\mu_b}{\sqrt d} +\log\!\left(\lVert H_b-\overline H\rVert_2+\varepsilon\right). \] (12c)

A block is kept exact when it has either high semantic relevance or a large deviation from the shared statistic. The corresponding error bound scales with the tail probability fraction and the maximum heterogeneity \(\max_{b\in U_i}\lVert H_b-\overline H\rVert_2\). The important conclusion is that selection and tail approximation must be designed together.

Recap: the same two questions, four answers

Only after seeing the mechanisms is the comparison useful. Each method must say both how it finds the exact set \(E_i\) and what becomes of the unselected set \(U_i\).

MethodHow \(E_i\) is selectedHow the selector learnsWhat happens to \(U_i\)
DSALow-dimensional per-token indexerKL distillation from dense attentionDropped; selected-set softmax
NSAScores from learned compressed blocksLM loss through a value-bearing compressed branchRaw KV is skipped; coarse information survives in that branch
HiLSLandmark estimate of chunk LogSumExpLM loss through selected inter-chunk massDropped; selected chunks use hierarchical softmax
PISACentroid relevance plus an approximation-error priorTraining-free statisticsApproximated in the same numerator and denominator

The selector must cost less than the tiles it skips

The cost model must include router FLOPs, top-\(k\), index storage, KV gathers, and the sparse kernel—not only the surviving dot products. A mathematically sparse mask can still be slow if its selected entries require irregular memory reads or produce tiny matrix multiplications.

Under grouped-query attention, multiple query heads share one physical KV head. If every query head selects unrelated blocks, the actual KV load becomes the union of those sets and the memory saving can disappear. This is why hardware alignment is part of the algorithmic contract, not a final implementation detail.

The reusable mental model First write attention as one numerator divided by one denominator. Then ask three questions: Which pairs are evaluated? What represents the omitted pairs? Can those choices be executed as large, regular blocks? This separates mathematical accuracy from kernel efficiency.

References

  1. [1]
    CUDA Refresher: The CUDA Programming Model Pradeep Gupta. NVIDIA Technical Blog, June 26, 2020.
  2. [2]
    FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness Tri Dao, Daniel Y. Fu, Stefano Ermon, Atri Rudra, and Christopher Ré. NeurIPS 2022.
  3. [3]
  4. [4]
    FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision Jay Shah, Ganesh Bikshandi, Ying Zhang, Vijay Thakkar, Pradeep Ramani, and Tri Dao. NeurIPS 2024.
  5. [5]
    Ring Attention with Blockwise Transformers for Near-Infinite Context Hao Liu, Matei Zaharia, and Pieter Abbeel. arXiv:2310.01889, 2023.
  6. [6]
    TileLang: A Composable Tiled Programming Model for AI Systems Lei Wang, Yu Cheng, Yining Shi, et al. arXiv:2504.17577, 2025.
  7. [7]
    Generating Long Sequences with Sparse Transformers Rewon Child, Scott Gray, Alec Radford, and Ilya Sutskever. arXiv:1904.10509, 2019.
  8. [8]
    Reformer: The Efficient Transformer Nikita Kitaev, Łukasz Kaiser, and Anselm Levskaya. ICLR 2020.
  9. [9]
    Longformer: The Long-Document Transformer Iz Beltagy, Matthew E. Peters, and Arman Cohan. arXiv:2004.05150, 2020.
  10. [10]
    Big Bird: Transformers for Longer Sequences Manzil Zaheer, Guru Guruganesh, Avinava Dubey, et al. NeurIPS 2020.
  11. [11]
    Efficient Content-Based Sparse Attention with Routing Transformers Aurko Roy, Mohammad Saffar, Ashish Vaswani, and David Grangier. Transactions of the Association for Computational Linguistics 9:53–68, 2021.
  12. [12]
    MInference 1.0: Accelerating Pre-filling for Long-Context LLMs via Dynamic Sparse Attention Huiqiang Jiang, Yucheng Li, Chengruidong Zhang, et al. NeurIPS 2024.
  13. [13]
    Quest: Query-Aware Sparsity for Efficient Long-Context LLM Inference Jiaming Tang, Yilong Zhao, Kan Zhu, Guangxuan Xiao, Baris Kasikci, and Song Han. ICML 2024.
  14. [14]
    SpargeAttention: Accurate and Training-free Sparse Attention Accelerating Any Model Inference Jintao Zhang, Chendong Xiang, Haofeng Huang, Jia Wei, Haocheng Xi, Jun Zhu, and Jianfei Chen. ICML 2025.
  15. [15]
    Native Sparse Attention: Hardware-Aligned and Natively Trainable Sparse Attention Jingyang Yuan, Huazuo Gao, Damai Dai, et al. ACL 2025.
  16. [16]
    MoBA: Mixture of Block Attention for Long-Context LLMs Enzhe Lu, Zhejun Jiang, Jingyuan Liu, et al. NeurIPS 2025.
  17. [17]
    Fast Video Generation with Sliding Tile Attention Peiyuan Zhang, Yongqi Chen, Runlong Su, Hangliang Ding, Ion Stoica, Zhengzhong Liu, and Hao Zhang. ICML 2025.
  18. [18]
  19. [19]
    Hierarchical Sparse Attention Done Right: Toward Infinite Context Modeling Xiang Hu, Xinyu Wei, Hao Gu, et al. arXiv:2607.02980, 2026.
  20. [20]
    PISA: Piecewise Sparse Attention Is Wiser for Efficient Diffusion Transformers Haopeng Li, Shitong Shao, Wenliang Zhong, Zikai Zhou, Lichen Bai, Hui Xiong, and Zeke Xie. arXiv:2602.01077, 2026.