Flattening Every Memory Peak in Long-Context Mixture-of-Experts Training
AuthorsShrey Pandit, Xuan-Phi Nguyen, Yiran Zhao, Shafiq Joty
Resources
This work makes trillion-scale, million-token MoE training more feasible by flattening every major GPU memory spike instead of merely reducing the average footprint.
Key results
Maximum memory reduction from PipelinedLLEP versus LLEP.
Maximum reduction from Ring-DTP with streamed vocabulary blocks.
OffloadStreamAdamW speedup over CPU AdamW for the offloaded optimizer step.
Context length reached by the composed system across tested MoE scales.
Largest MoE model trained end to end with the composed stack.
Maximum throughput relative to a tuned FSDP2 baseline.
What the paper found
This paper tackles a failure mode in long-context Mixture-of-Experts training: preventing every individual memory peak from exceeding GPU capacity, rather than merely reducing average usage. It identifies four independently growing live sets—MoE dispatch buffers, vocabulary logits, gradient-checkpoint boundaries, and AdamW state—and introduces four exact streaming schedules: PipelinedLLEP caps per-source tokens per dispatch chunk; Ring-DTP circulates activations or vocabulary shards while accumulating cross-entropy with online log-sum-exp; Selective Checkpoint Offload prefetches checkpoint boundaries from host memory; and OffloadStreamAdamW pipelines optimizer buckets through NVIDIA H200 GPUs. The methods preserve the model, loss, gradients, and BF16 training semantics, unlike quantization or approximate routing. On controlled tests, PipelinedLLEP cuts dispatch peak memory by up to 59.3%, Ring-DTP cuts vocabulary-projection memory by 86.6% with under 5% overhead, and OffloadStreamAdamW makes the offloaded optimizer step 2.05× faster. Combined with a Mixture-of-Parallelisms layout, the system trains 120B-to-667B-parameter MoE models at 1M-token context, reaching up to 10.4× the throughput of tuned FSDP2 baselines. Experiments use OpenAI’s gpt-oss-20b and DeepSeek-V3-style auxiliary-loss-free routing, demonstrating that memory savings can extend context or batch size without changing training quality.
Original abstract
Training a Mixture-of-Experts (MoE) model at long context or large batch size fails as soon as any one component's peak allocation exceeds device memory, so the target is every peak at once, not the average footprint. Four are left unbounded by the parallelism plans in common use, and each grows differently: expert dispatch with the routing matrix, the vocabulary projection with tokens times vocabulary, gradient checkpoint boundaries with depth times sequence length, and optimizer state with parameter count. Which one runs out first changes with the model, the context length, and the device count, so lowering the largest only exposes the next. We bound all four with schedules whose GPU working set is fixed at launch: PipelinedLLEP extends least-loaded expert parallelism with a cap on the tokens each source contributes to a dispatch chunk, Ring-DTP circulates activations or weight shards around a ring at the vocabulary projection and folds each block of logits into an online log-sum-exp, Selective checkpoint offload (SCO) keeps the one long-lived tensor of each checkpoint boundary in CPU memory, and OffloadStreamAdamW turns the serial CPU Adam update of optimizer offload into a bucket pipeline. All four change only the order and granularity of computation and data movement, so the loss and gradients stay exact. In matched component tests, they cut the MoE dispatch peak by up to $59.3\%$ without losing throughput, the vocabulary projection peak by $86.6\%$, and the offloaded optimizer step by $2.05\times$ faster. Composed on MoE models from 120B to 667B parameters, they train at 1M context length, $8$--$32\times$ the reach of a tuned FSDP2 baseline, and up to $10.4\times$ its throughput.
Read the original paperMore in Efficient AI
Browse all 55 papers →Decoding Looped Transformers Better for (Almost) Free
Weihao Liu, Huangjie Zheng, Tianrong Chen, Rohit Dilip, Richard He Bai, Yizhu Jiao, Yuyang Wang, Ruixiang Zhang
LoopCD turns the partially computed states of looped Transformers into free guidance, improving accuracy while often cutting inference compute nearly in half.
Scaling Laws for Looped Mixture of Experts
Yanbei Chen, Anirudh Goyal, Raghuraman Krishnamoorthi
This work develops scaling laws that explain how looping and sparse experts can be combined to build more capable models with less training and inference compute.
When Fancy Eviction Fails: Rethinking Cache Replacement For LLM Prefix Reuse
Yiyu Liu, Minlan Yu, Juncheng Yang
For LLM prefix caches, simple recency may beat fancy eviction rules, especially when workloads follow predictable session patterns.