cs.LGSep 27, 2026

FoldAttention: Declared-Reference Softmax for Fast Decode and Deterministic Backward

Authors: Sriman Achanta

Organizations: Virginia Commonwealth University

Abstract

Autoregressive decode repeatedly streams a growing KV cache, making attention a major cost at long context. Existing high-performance kernels use online softmax, which discovers a row's normalization reference as it scans keys. Earlier contributions therefore remain provisional and may require rescaling. We argue that the reference need not be discovered: softmax is invariant to a common shift, so the reference only has to keep the weights in range. We present FoldAttention, an additive formulation of softmax attention that fixes a finite reference ZiZ_i before scanning the KV cache. Each weight 2sij−Zi2^{s_{ij}-Z_i} is then final when computed, so contributions add across disjoint key ranges and their quotient equals softmax attention in real arithmetic. We use this property to develop two techniques for Hopper decode: (1) final weights gate key and value reads before the bytes are fetched, and a per-call depth TT cuts keys below 2−T2^{-T} while keeping their mass, and (2) additive partials compose split KV and shared-prefix cascades without rescaling. On H100 at T=16T=16, FoldAttention decodes seven real-model generations 1.36-2.30×\times faster than the fastest BF16 baseline, and up to 3.09×\times faster across MHA and GQA shapes, at an error within 1.5% of the lowest BF16 error on six of the seven; reading every key, it is 1.14-1.30×\times faster at matched error. We validate on Qwen3-8B that a whole decode step is up to 1.46×\times faster while likelihood and long-context accuracy match those under BF16 kernels. The same principle makes the backward deterministic: CTAs round bounded partial gradients onto an integer grid declared before the reduction and add them in any order. FoldAttention thereby removes the determinism tax: its deterministic backward is up to 1.84×\times faster than deterministic FlashAttention-3/4 and 1.05×\times faster than the fastest nondeterministic kernel.

Figures & tables

Appendix figures & tables5 assets

Supplementary material from the paper’s appendix.

Appendix

Explore similar work

Sep 21, 2026cs.LG

FlashBoB: I/O-Efficient Exact Backward-over-Backward for Softmax Attention

Transformer models built on the attention mechanism have become a central building block in modern deep learning, yet softmax attention remains a major bottleneck for long-context workloads. While FlashAttention makes the forward and first backward passes I/O-efficient, it does not support backward-over-backward (BoB), which enables exact differentiation through the backward pass for applications such as second-order optimization, test-time training, gradient-based memory, and meta-learning. Existing BoB implementations either materialize large intermediate tensors or exhaust GPU memory at long sequence lengths. We present FlashBoB, an exact, I/O-efficient algorithm for BoB in softmax attention that keeps computation within on-chip tiles and avoids all N×NN \times N intermediate tensors, where NN is the sequence length. The key insight is a hierarchical affine structure in the softmax double backward: two row-wise scalars determine all outputs through affine transformations. This yields a two-pass schedule with bounded on-chip static random-access memory (SRAM) usage and minimal off-chip high-bandwidth memory (HBM) traffic. FlashBoB achieves Θ(N2d2/M)Θ(N^2 d^2/M) HBM traffic (dd is the head dimension and MM is the memory size) and, within the standard FlashAttention-style score-recomputation model, matches the inherited large-cache lower bound for exact forward attention. Empirically, it scales exact attention BoB to N=262KN=262\text{K} on a single A100 80GB GPU, where prior PyTorch exact baselines fail by N=16KN=16\text{K}, and is up to 6.3×6.3\times faster than FlashBack. These results make exact second-order attention practical at long-context sequence lengths where prior implementations cannot run efficiently.
May 3, 2026cs.LG

Stochastic Sparse Attention for Memory-Bound Inference

Autoregressive decoding becomes bandwidth-limited at long contexts, as generating each token requires reading all nkn_k key and value vectors from KV cache. We present Stochastic Additive No-mulT Attention (SANTA), a method that sparsifies value-cache access by sampling S≪nkS \ll n_k indices from the post-softmax distribution and aggregates only those value rows. This yields an unbiased estimator of the post-softmax value aggregation while replacing value-stage multiply-accumulates with gather-and-add. We introduce stratified and systematic sampling to design variance-reduced, GPU-friendly variants. Evaluated on Llama-3.1-8B-Instruct at 32k-token contexts, S2^2ANTA matches baseline accuracy while achieving up to 1.5×1.5\times decode-step attention-kernel speedup over FlashInfer and FlashDecoding on an NVIDIA RTX 6000 Ada. In batched long-context generation, these kernel gains translate to up to 1.25×1.25\times end-to-end decode-latency speedup. Finally, we propose Bernoulli qKTqK^\mathsf{T} sampling as a complementary technique to sparsify the score stage, reducing key-feature access through stochastic ternary queries. Both methods are complementary to upstream quantization, low-rank projection, KV-cache compression, and KV-cache selection methods. Together, they point toward sparse, multiplier-free, and energy-efficient inference. We open-source our kernels at: https://github.com/OPUSLab/SANTA.git
Dec 18, 2025cs.LG

Kascade: A Practical Sparse Attention Method for Long-Context LLM Inference

Attention is the dominant source of latency during long-context LLM inference, an increasingly popular workload with reasoning models and RAG. We propose Kascade, a training-free sparse attention method that leverages known observations such as 1) post-softmax attention is intrinsically sparse, and 2) the identity of high-weight keys is stable across nearby layers. Kascade computes exact Top-k indices in a small set of anchor layers, then reuses those indices in intermediate reuse layers. The anchor layers are selected algorithmically, via a dynamic-programming objective that maximizes cross-layer similarity over a development set, allowing easy deployment across models. The method incorporates efficient implementation constraints (e.g. tile-level operations), across both prefill and decode attention. The Top-k selection and reuse in Kascade is head-aware and we show in our experiments that this is critical for high accuracy. Kascade achieves up to 4.1x speedup in decode attention and 2.2x speedup in prefill attention over FlashAttention-3 baseline on H100 GPUs while closely matching dense attention accuracy on long-context benchmarks such as LongBench and AIME-24.