Skip to content

Block Sparse Flash Attention

Conference: NeurIPS2026
arXiv: 2512.07011
Code: https://github.com/Danielohayon/Block-Sparse-Flash-Attention
Area: LLM Efficiency
Keywords: block-sparse attention, exact-score gating, threshold calibration, long context, prefill acceleration

TL;DR

BSFA first computes all causally visible QK scores exactly inside FlashAttention-2, then uses offline-calibrated block-maximum thresholds to skip V loads, PV, and softmax-state updates for low-scoring blocks, achieving a LongBench score of 39.78% versus the dense baseline's 40.24% on Llama-3.1-8B and a 1.13ร— end-to-end prefill speedup on the longest 10 samples.

Background & Motivation

The first output token in long-context inference must wait for prefill of the entire input. FlashAttention-2 uses tiling and online softmax to avoid writing the full attention matrix to device memory, but every causally visible block still requires QK, exponentiation and normalization, and PV. QK determines relevance, whereas PV adds the corresponding content to the output; their main matrix multiplications have the same size, so removing attention-matrix storage does not remove arithmetic that grows quadratically with context length. In long-document question answering, many positions ultimately receive negligible attention weight, leaving room to skip some value-side work.

Many sparse methods decide which blocks deserve computation before observing their true scores. MInference searches sparse patterns, while FlexPrefill and XAttention select important regions using local or compressed score approximations. This can save both QK and PV, but it can also miss distant, scattered information or signals poorly represented by the proxy; the extra selection procedure can cost more than it saves on short contexts. Rather than improving a cheaper predictor, BSFA retains all QK computation and uses the resulting true scores to decide whether the second half is worth executing. It cannot remove QK's quadratic complexity, but selection has more reliable information and gating fits directly into the original kernel.

A deployment issue remains: sorting every block outside the kernel for exact top-k selection on each input could offset sparsity gains. The authors exploit relatively stable score distributions across layers, heads, and positions within a model, calibrating lookup thresholds in advance to replace online sorting with one comparison. Core Idea: pay the full QK cost to obtain exact block maxima, then approximate a target block budget using thresholds calibrated by layer, head, and position, continuing value aggregation and online softmax only for retained blocks.

Method

Overall Architecture

BSFA is a prefill attention replacement kernel that requires no model-weight updates. Offline calibration builds threshold tables for several block budgets; at inference time, Q, K, V, and the threshold slice for the selected budget produce an attention output renormalized over retained blocks. It neither constructs the full score matrix in global device memory nor permanently deletes tokens from the KV cache.

The stages are โ€œOffline Threshold Calibration,โ€ โ€œExact-Score Gating,โ€ and โ€œSelective Streaming Aggregationโ€: the first supplies deployment parameters, and the latter two execute inside each attention call. The dashed edge carries only an offline calibration artifact, not training gradients or an additional online selection network; the causal diagonal region proceeds directly to aggregation, while off-diagonal regions are retained or skipped by the gate.

%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400}}}%%
flowchart TD
    C["Calibration samples"] --> T["Offline Threshold<br/>Calibration"]
    X["Inference Q, K, V"] --> G["Exact-Score Gating"]
    T -.->|Threshold table| G
    G -->|Retained blocks and diagonal region| A["Selective Streaming<br/>Aggregation"]
    G -->|Skipped block: unchanged state| N["Next key block"]
    A --> N
    N -->|More key blocks| G
    N -->|Traversal complete| O["Normalized output"]

Key Designs

1. Offline Threshold Calibration: replace online top-k sorting with a lookup comparison

For every calibration sample, layer, attention head, and query-block position, the method first computes the maximum score against every visible off-diagonal key block, then sorts these maxima and takes the kth highest score as that sample's threshold. Averaging these thresholds across calibration samples yields the deployment table. Layers and heads differ in their preferences for local context, initial attention-sink positions, or remote content, and query position determines the amount of visible history; a single model-wide threshold would struggle to capture these differences.

This procedure uses the target number of off-diagonal blocks \(k\) as a budget rather than imposing a uniform sparsity ratio across the sequence. The authors use 16 RULER calibration samples and separate their task categories from those used for latency evaluation; the same thresholds are used on LongBench instead of retuning for LongBench. Multiple budgets can be stored as separate slices, allowing deployment-time switching between speed and quality settings without rerunning calibration.

The calibration target must be distinguished from actual execution: an averaged threshold is not the kth largest score of a new input, so more or fewer than k blocks can exceed it. The paper's โ€œfixed-kโ€ should therefore be read as a target-budget design, not a guarantee that every input retains exactly k blocks; its intent to reduce workload variation must not be presented as completely eliminating thread-level imbalance.

For early query positions with insufficient off-diagonal candidates, the threshold is \(-\infty\), avoiding forced pruning of short histories. The appendix also reuses the final calibrated position's threshold beyond the maximum calibration length; this is an empirical extrapolation strategy, not a proof that distributions remain unchanged at arbitrary lengths. Changing the model or tile shape also changes the threshold's meaning, so it cannot be reused indiscriminately.

2. Exact-Score Gating: observe the full score block before deciding value-side work

At inference time, the query block stays on-chip while the kernel loads keys block by block, computes scaled dot-product scores, and takes the maximum over the entire queryโ€“key score tile. This maximum covers all query and key positions inside the block; it is not independent token-level top-k selection. One sufficiently relevant pair can retain the block. The gate uses pre-softmax scores rather than normalized attention probabilities.

\[ s_{\max}^{(i,j)}=\max_{p,q}\left[\frac{\mathbf Q_i\mathbf K_j^{\top}}{\sqrt d}\right]_{pq},\qquad \text{keep}(i,j)\iff s_{\max}^{(i,j)}\ge T_{\ell,h,i}^{(k)}. \]

Here, \(i,j\) denote conceptual query-block and key-block positions, and \(\ell,h\) identify the layer and head. โ€œExactโ€ means QK scores are not approximated by mean pooling, sampling, or another model; it does not mean that the entire sparse output is exactly equal to dense attention. All causally visible QK blocks are still computed; the gate saves subsequent value-side processing rather than avoiding key-side scoring in advance.

The causal diagonal region always executes, with a causal mask applied to future positions, preserving current local context. The main text and Figure 2 use โ€œgreater than or equal to,โ€ whereas simplified Algorithm 1 uses strict โ€œgreater thanโ€; this note follows the main text, and equality handling should be checked against the implementation. The algorithm also uses simplified diagonal-block indexing, while actual A100 query and key tile sizes differ, so its single \(j=i\) index is not a complete mapping for unequal CUDA tile shapes.

Block maxima are conservative for a specific reason: a block is dropped only if it contains no high-scoring pair, protecting scattered relevant tokens. However, one exceptionally high-scoring query can retain a block for the other queries in the tile; conversely, many individually low-scoring blocks whose combined contribution is useful can still be removed. The maximum supplies a reliable scoring signal, not a rigorous error bound on discarded probability mass.

3. Selective Streaming Aggregation: skipped blocks neither load V nor enter the denominator

For a retained block, the kernel performs FlashAttention-style online softmax: update the per-query running maximum, compute stable exponentials, rescale the accumulated output and normalizer when needed, then load V and accumulate PV. After traversal, the accumulated normalizer scales the output. Score tiles, state, and output accumulators remain on-chip rather than being written as a full attention matrix.

For a rejected block, the kernel does not load its V, execute PV, compute its exponentials, or update the running maximum, normalizer, and output accumulator. This is stronger than simply setting the block's PV contribution to zero: its contribution to the softmax denominator is removed as well. The final result is approximate attention renormalized over the retained block set, not exact value aggregation under dense softmax probabilities over all keys.

This explains why model quality still requires empirical evaluation. If low-scoring blocks carry little probability mass, removing them usually causes a small normalization perturbation; exact QK alone does not prove that their total discarded mass is zero. The paper relies on empirical calibration and task accuracy, not a lossless guarantee for every input.

For each rejected tile, the main QK and PV matrix multiplications each require \(2B_MB_Nd\) FLOPs, so skipping PV saves half their combined cost, with additional softmax savings. K and V have the same size, so avoiding V saves 50% of that tile's KV reads. This 50% is a local accounting limit for KV traffic or the two main matrix multiplications in a pruned tile, not a 50% reduction for the entire layer, all attention, or the full inference pipeline; retained blocks, all QK, projections, MLPs, and kernel scheduling remain.

A Worked Example

The following hypothetical example explains the mechanism and is not experimental data from the paper. Suppose a query block can see 5 off-diagonal candidate blocks, the deployment budget is \(k=2\), the calibrated threshold is 3, and this input's block maxima are 1, 4, 2, 5, and 3.5.

The kernel still loads all 5 key blocks and computes their QK, then skips the two blocks with maxima 1 and 2 and retains the three value blocks corresponding to 4, 5, and 3.5. Retaining 3 rather than 2 blocks demonstrates that an offline-averaged threshold only approximates the budget. The causal diagonal region is retained additionally and does not consume the off-diagonal budget.

When the block scoring 2 is skipped, the previously accumulated maximum, softmax denominator, and weighted output remain unchanged; if the later block scoring 5 raises the running maximum, the existing state is stably rescaled. Final normalization covers only the three retained off-diagonal blocks and the causal diagonal region, excluding skipped blocks from the denominator. If the query position is too early to supply enough candidates, a \(-\infty\) threshold retains the available history.

Loss & Training

BSFA introduces no loss function, trains no selector, and does not fine-tune LLM weights. Offline calibration estimates thresholds of score distributions rather than a supervised learning objective; deployment uses the original model's QK and stored thresholds for selection.

The main experiments use an A100 80GB, CUDA 12.1, FP16, query tiles of 128, and key tiles of 64. The H100 port changes the query tile to 128 and the key tile to 224 and recalibrates because tile size changes; it uses an FA-2-style algorithm, not completed FA-3 integration. A6000 reuses A100 thresholds under the same tile semantics, whereas Qwen2.5-7B is calibrated separately; this does not imply that different models can share thresholds.

Key Experimental Results

Main Results

Accuracy in this table comes from the full LongBench benchmark of 27 tasks and approximately 4,750 samples; speedup comes from the longest 10 narrativeqa samples, each approximately 65K tokens. The columns do not describe the same evaluation subset. Unless stated otherwise, the model is Llama-3.1-8B on A100 and the baseline is Dense FlashAttention-2; degradation is relative to baseline rather than a percentage-point difference.

Method Parameter LongBench accuracy Relative degradation End-to-end TTFT speedup
Dense FlashAttention-2 All 40.24% โ€” 1.00ร—
BSFA k=32 39.39% 2.1% 1.16ร—
BSFA k=64 39.78% 1.1% 1.13ร—
BSFA k=96 39.88% 0.9% 1.10ร—
BSFA k=128 40.03% 0.5% 1.08ร—
MInference default 40.06% 0.45% 1.17ร—
XAttention default 39.95% 0.7% 1.16ร—
FlexPrefill ฮณ=0.99 38.90% 3.3% 1.06ร—

These results support BSFA's adjustable qualityโ€“latency settings, but not the claim that it is the only method reaching 99% relative accuracy in Table 1: MInference and XAttention also reach that range in this table. BSFA's additional advantage requires considering regressions on short and medium inputs. BLASST's 50%-sparsity LongBench configuration scores 39.23% at 1.11ร—, but runs on H100 and should not be folded into this table as a strict same-GPU ranking.

Ablation Study

The paper mainly analyzes budgets and lengths rather than removing individual modules. This table combines kernel timings from original Table 3 with density data from Appendix Table 6, all for RULER 64K, Llama-3.1-8B, and A100. Predicted density follows the target budget, while measured density is the actual retained fraction; their difference also directly demonstrates that threshold gating does not guarantee exact k.

Configuration RULER accuracy Predicted / measured density End-to-end speedup Attention-kernel speedup
Dense reference 84.96% 1.00 / 1.00 1.00ร— 1.00ร—
BSFA k=192 83.08% 0.35 / 0.36ยฑ0.05 1.13ร— 1.38ร—
BSFA k=256 83.39% 0.44 / 0.45ยฑ0.05 1.09ร— 1.30ร—
BSFA k=384 84.24% 0.62 / 0.61ยฑ0.05 1.05ร— 1.19ร—

Original Table 3 labels its dense row SDPA, whereas Appendix Table 6 labels it Dense FlashAttention-2; this note retains โ€œdense referenceโ€ without assuming that the two implementation labels are identical. Kernel timing excludes Q/K/V/O projections, RoPE, GQA expansion, MLPs, and the LM head, so 1.38ร— must not be reported as a 1.38ร— end-to-end gain.

A separate analysis uses an independent mixed-length LongBench subset of 160 samples spanning 295โ€“65,461 tokens. The following table reports only end-to-end TTFT speedup and does not treat full-LongBench accuracy as per-bucket accuracy.

Method <4K 4โ€“8K 8โ€“16K 16โ€“32K 32โ€“65K
BSFA k=64 0.96ร— 0.97ร— 1.01ร— 1.05ร— 1.19ร—
MInference default 0.17ร— 0.31ร— 0.48ร— 0.58ร— 0.99ร—
FlexPrefill ฮณ=0.99 0.69ร— 0.80ร— 0.93ร— 0.99ร— 1.12ร—
XAttention default 0.85ร— 0.90ร— 0.97ร— 1.01ร— 1.09ร—

Key Findings

  • Smaller block budgets improve speed, but RULER 64K at k=192 still falls 1.88 percentage points below the dense score. Exact QK does not imply lossless quality.
  • BSFA also regresses slightly below 8K rather than accelerating every length; its advantage is staying closer to dense latency than baselines with additional online block selection. The 1.13ร— result on the longest 10 samples and the 1.19ร— result in the mixed subset's 32โ€“65K bucket are not interchangeable.
  • Needle-in-a-Haystack 64K single-key retrieval reaches 99% of baseline accuracy and 1.24ร— speedup at k=32, which does not generalize to all tasks requiring distributed evidence; the other four major sparse baselines were not measured on this task.

Highlights & Insights

  • Retain scoring, prune the work consuming scores: relevance need not be predicted before it is observed, shifting selection risk to whether low-scoring blocks are negligible. This retains an exact scoring signal but also imposes an honest, lower ceiling on savings.
  • Calibrate a budget, not a fixed window: layer-, head-, and position-specific thresholds can select remote content, whereas the same-budget k=64 sliding window scores only 13.43% on LongBench. Actual block counts remain variable after threshold averaging, so deployment should monitor density rather than recording k alone.

Limitations & Future Work

  • Retaining all QK keeps arithmetic complexity quadratic; gains are limited jointly by skippable PV and non-attention overhead, so this is not a subquadratic attention method.
  • Thresholds depend on the model, position, and tile shape. Cross-dataset transfer has empirical support, but terminal-threshold reuse at longer positions, distribution shifts, and accumulated low-score mass lack rigorous error guarantees.
  • The method targets prefill and does not establish decode acceleration or KV-cache capacity reduction. H100 results use an FA-2-style baseline and do not establish superiority over FA-3; compatibility with FA-3 scheduling or sparse-output correction remains a possible direction.
  • Baseline implementations are not fully homogeneous: SpargeAttention uses INT8 Q/K quantization, FlexPrefill uses bf16 with a dtype-matched dense baseline, and BLASST uses H100. These conditions limit cross-method attribution of speed and quality differences.
  • vs FlashAttention-2: FA-2 processes all visible blocks exactly; BSFA retains its tiling and online normalization but excludes some blocks, making it an approximate attention replacement rather than another lossless IO optimization.
  • vs BLASST: both gate after scoring; BSFA uses absolute thresholds calibrated by layer, head, and position, while BLASST gates relative to a running maximum. Although the paper describes BSFA as fixed-budget, execution is still threshold-driven, so the distinction must not be absolutized as โ€œfixed versus variable.โ€
  • vs MInference / FlexPrefill / XAttention: these restrict QK using patterns or proxy scores first and can theoretically save more computation; BSFA saves less QK in exchange for observing full scores and lower online selection overhead. Choosing between these approaches depends on length and quality requirements.
  • vs KV-cache compression / ฮ”-Attention: cache compression addresses capacity or decode bandwidth, and ฮ”-Attention corrects sparse-output shifts, both distinct from BSFA's prefill value-block selection; potential combinations require independent validation.

Rating

  • Novelty: 4/5 โ€” Combines exact block scores with fine-grained offline thresholds; the contribution is clear but adjacent to existing post-score gating.
  • Experimental Thoroughness: 4/5 โ€” Covers long-document tasks, budgets, lengths, models, and devices, but separated accuracy/timing subsets, cross-device baselines, and implementation differences remain relevant.
  • Writing Quality: 3/5 โ€” The core algorithm is intuitive, but fixed-k wording, threshold equality handling, and dense-reference labels are inconsistent.
  • Value: 4/5 โ€” Useful for quality-sensitive FA-2 prefill deployments with variable input lengths; gains are practical rather than orders of magnitude.