Paged Attention Meets FlexAttention: Unlocking Long-Context Efficiency in Deployed Inference

📅 2025-06-08
📈 Citations: 0
✨ Influential: 0
📄 PDF
🤖 AI Summary
To address memory inefficiency and internal fragmentation in traditional KV caching for long-context reasoning in large language models (LLMs), this work proposes a fused attention mechanism integrating PagedAttention with PyTorch FlexAttention. Our novel design enables dynamic aggregation of disjoint KV blocks and supports non-contiguous memory layouts, eliminating fragmentation inherent in monolithic cache structures and enabling low-overhead long-sequence inference. Leveraging a custom CUDA fused kernel and integrated within the IBM Foundation Model Stack (FMS), experimental evaluation on an NVIDIA L4 GPU demonstrates near-linear latency scaling—approximately 2× increase—for sequences of 128–2048 tokens, while peak memory consumption remains nearly constant. A marginal increase in memory usage emerges only beyond 2048 tokens, attributable to power-of-two paging granularity.

Technology Category

Machine Learning: Large Multimodal Models (LMMs)Natural Language Processing: (Large) Language ModelsComputer Vision: Large Vision Models

Application Category

Semantics and Knowledge: Methods to enhance, augment, integrate or synergize semantic models such as knowledge graphs and LLMsSearch and Retrieval-Augmented AI: Large language models for searchGraph Algorithms and Modeling for the Web: Foundation models and LLMs for Web-related graphs
📝 Abstract
Large Language Models (LLMs) encounter severe memory inefficiencies during long-context inference due to conventional handling of key-value (KV) caches. In this work, we introduce a novel integration of PagedAttention with PyTorch's FlexAttention, addressing internal fragmentation and inefficiencies associated with monolithic KV cache allocations. Implemented within IBM's Foundation Model Stack (FMS), our fused attention kernel efficiently gathers scattered KV data. Our benchmarks on an NVIDIA L4 GPU (24GB) demonstrate significantly reduced inference latency, growing only linearly (~2x) with sequence length from 128 to 2048 tokens when utilizing a global KV cache, compared to exponential latency increases without caching. While peak memory usage remains largely unchanged for single-step evaluations (dominated by model weights and activations), paged attention causes minimal incremental memory usage, observable only at sequence lengths exceeding 2048 tokens due to its power-of-two cache allocations. We open-source the full implementation and discuss its implications for future long-context model deployment.
Problem

Research questions and friction points this paper is trying to address.

Addresses memory inefficiencies in long-context LLM inference
Integrates PagedAttention with FlexAttention to reduce fragmentation
Improves linear latency scaling with global KV caching
Innovation

Methods, ideas, or system contributions that make the work stand out.

Integrates PagedAttention with PyTorch's FlexAttention
Efficiently gathers scattered KV cache data
Reduces inference latency linearly with sequence length
🔎 Similar Papers
No similar papers found.
💼 Related Jobs
No related jobs found.
T
Thomas Joshi
Columbia University
H
Herman Saini
Columbia University
N
Neil Dhillon
Columbia University
A
Antoni Viros i Martin
IBM
K
K. E. Maghraoui
IBM