🤖 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.
📝 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.