Table of Contents
Fetching ...

Efficient Long-context Language Model Training by Core Attention Disaggregation

Yonghao Zhuang, Junda Chen, Bo Pang, Yi Gu, Yibo Zhu, Yimin Jiang, Ion Stoica, Eric Xing, Hao Zhang

TL;DR

This work tackles the bottleneck of load imbalance in long-context LLM training caused by the quadratic complexity of core attention ($O(l^2)$) relative to linear non-attention components. It introduces core attention disaggregation (CAD), which isolates core attention on dedicated attention servers and leverages token-level shard fusion, a ping-pong execution scheme, and a communication-aware greedy scheduler implemented in DistCA. The approach yields up to 1.35x end-to-end throughput improvements on 512 H200 GPUs for context lengths up to 512k tokens and eliminates data and pipeline stragglers, demonstrating near-linear balance at scale. The method enables efficient long-context training with practical integration into existing pipelines and suggests broad applicability to future large-scale LLMs.

Abstract

We present core attention disaggregation (CAD), a technique that improves long-context large language model training by decoupling the core attention computation, softmax(QK^T)V, from the rest of the model and executing it on a separate pool of devices. In existing systems, core attention is colocated with other layers; at long context lengths, its quadratic compute growth compared to the near-linear growth of other components causes load imbalance and stragglers across data and pipeline parallel groups. CAD is enabled by two observations. First, core attention is stateless: it has no trainable parameters and only minimal transient data, so balancing reduces to scheduling compute-bound tasks. Second, it is composable: modern attention kernels retain high efficiency when processing fused batches of token-level shards with arbitrary lengths. CAD partitions core attention into token-level tasks and dispatches them to dedicated attention servers, which dynamically rebatch tasks to equalize compute without sacrificing kernel efficiency. We implement CAD in a system called DistCA, which uses a ping-pong execution scheme to fully overlap communication with computation and in-place execution on attention servers to reduce memory use. On 512 H200 GPUs and context lengths up to 512k tokens, DistCA improves end-to-end training throughput by up to 1.35x, eliminates data and pipeline parallel stragglers, and achieves near-perfect compute and memory balance.

Efficient Long-context Language Model Training by Core Attention Disaggregation

TL;DR

This work tackles the bottleneck of load imbalance in long-context LLM training caused by the quadratic complexity of core attention () relative to linear non-attention components. It introduces core attention disaggregation (CAD), which isolates core attention on dedicated attention servers and leverages token-level shard fusion, a ping-pong execution scheme, and a communication-aware greedy scheduler implemented in DistCA. The approach yields up to 1.35x end-to-end throughput improvements on 512 H200 GPUs for context lengths up to 512k tokens and eliminates data and pipeline stragglers, demonstrating near-linear balance at scale. The method enables efficient long-context training with practical integration into existing pipelines and suggests broad applicability to future large-scale LLMs.

Abstract

We present core attention disaggregation (CAD), a technique that improves long-context large language model training by decoupling the core attention computation, softmax(QK^T)V, from the rest of the model and executing it on a separate pool of devices. In existing systems, core attention is colocated with other layers; at long context lengths, its quadratic compute growth compared to the near-linear growth of other components causes load imbalance and stragglers across data and pipeline parallel groups. CAD is enabled by two observations. First, core attention is stateless: it has no trainable parameters and only minimal transient data, so balancing reduces to scheduling compute-bound tasks. Second, it is composable: modern attention kernels retain high efficiency when processing fused batches of token-level shards with arbitrary lengths. CAD partitions core attention into token-level tasks and dispatches them to dedicated attention servers, which dynamically rebatch tasks to equalize compute without sacrificing kernel efficiency. We implement CAD in a system called DistCA, which uses a ping-pong execution scheme to fully overlap communication with computation and in-place execution on attention servers to reduce memory use. On 512 H200 GPUs and context lengths up to 512k tokens, DistCA improves end-to-end training throughput by up to 1.35x, eliminates data and pipeline parallel stragglers, and achieves near-perfect compute and memory balance.
Paper Structure (21 sections, 9 equations, 12 figures, 5 tables)

This paper contains 21 sections, 9 equations, 12 figures, 5 tables.

Figures (12)

  • Figure 1: Transformer and its workload imbalance caused by core attention.
  • Figure 2: DistCA architecture.
  • Figure 3: Latency and memory breakdown for all-gather in Context Parallel, Llama-8B. Document lengths are all 32k.
  • Figure 4: Throughput and memory divergence for variable-length data chunks, under different parallelism for 512K-token data chunks on 8B model.
  • Figure 5: Throughput of core attention.
  • ...and 7 more figures