Making FlashAttention-4 faster for inference
Blog post from Modal
Modal describes a series of contributions to FlashAttention-4 aimed at improving large language model inference, particularly memory-bandwidth-bound token decoding workloads with variable batch sizes, sequence lengths, and paged KV caches. The work focuses on changing parallelism from query-centric execution toward key/value splitting, adding support for irregular memory access through cp.async rather than TMA, and using CuTe DSL to develop specialized kernel variants. Added FP8 attention-input support improved throughput by up to 1.16× while reducing KV-cache memory requirements, while arbitrary KV page-size support improved compatibility and cache efficiency; a follow-up address-generation optimization raised small-page performance by up to 2.40×. Porting split-KV, or Flash-Decoding, increased throughput by up to 4.37× for small query lengths by distributing a query’s KV work across multiple GPU multiprocessors, though it requires a reduction kernel and can introduce small floating-point differences. Other changes reduce unnecessary work for short query sequences, improving single-token decode throughput by up to 3.06×, and extend grouped-query attention packing to irregular head ratios, producing a reported 2.92× improvement for one decode benchmark. The authors argue that flexible tile-level programming models and adaptive choices between memory-access and parallelization strategies are central to future high-performance attention kernels.
Use this post, company, and trend context to find content marketing opportunities, perform competitive analysis, or address product feature gaps via the Plushcap MCP server or the Plushcap API.