GPU Hierarchy Meets Structured Matrices: Fast Algorithms for State-Space Models
Berlin Chen ⋅ Caitlin Wang ⋅ Aakash Sunil Lahoti ⋅ Kevin Li ⋅ Jay Shah ⋅ Jack Carlisle ⋅ Timmy Liu ⋅ Mengyu Guo ⋅ Zico Kolter ⋅ Albert Gu ⋅ Tri Dao
Abstract
Linear-time sequence models such as State Space Models (SSMs) offer an efficient alternative to self-attention and have demonstrated strong language modeling performance at scale. In practice, however, their asymptotic advantage is often lost at common training and prefill lengths---optimized FlashAttention kernels remain faster than existing SSM implementations up to 8k tokens. In this work, we substantially narrow this performance gap with novel algorithmic improvements: First, we introduce a split-sequence algorithm that decouples the compute tile from the sequence-parallel split length, preserving efficient matrix multiplication tiles while substantially reducing boundary-state traffic and latency. Second, we resolve a longstanding *speed--stability tradeoff* in computing the 1-semiseparable (SS) decay mask, the main non-matmul bottleneck shared by all chunk-wise linear-time kernels (Mamba-2, Mamba-3, GDN, KDA). Existing kernels use the fast *diff segsum* (prefix subtraction), which suffers monotonicity violations and catastrophic cancellation---failures observed in real pretraining (e.g., Nemotron-H); the stable alternative (*direct segsum*) materializes the dense mask and is $\sim30\\%$ slower. Our *hierarchy-aware segsum* decomposes the SS mask into 1-SS diagonal blocks and rank-one off-diagonal blocks, computed via warp-level scans and cross-warp aggregation directly in registers, achieving the stability of direct segsum at the speed of diff segsum. We implement these algorithms using CuTe-DSL with warp specialization, asynchronous TMA, WGMMA, and layout-aware accumulator fragments. Across a large range of sequence length, our H100 kernels reach 91\% of measured memcpy bandwidth, improve prefill time by $4.1\times$ over the official Triton Mamba-2 kernel, and **outperform the highly-optimized FlashAttention-4 starting at 1-2k sequence length**---to our knowledge, the first time a linear-time sequence layer is faster than state-of-the-art attention at the 2-4k sequence lengths used in common pretraining and prefill workloads, where prior linear-time kernels only reached parity at $\geq 8$k. Our hierarchy-aware segsum is $1.35$-$1.40\times$ faster than a naive stable direct-segsum implementation.
Chat is not available.
Successful Page Load