Gated All-Reduce: Faster Inference with Post-Training Communication Pruning
Abstract
While tensor-parallelism reduces overall latency during large language models' inference, it introduces two all-reduce operations in every transformer block, which account for a large share of inference latency. We propose Gated All-Reduce, a post-training method that removes a predefined fraction of these all-reduce operations. When skipping an all-reduce operation, each device accumulates its local contribution and merges it at the next retained all-reduce, so every layer's output eventually reaches the shared residual stream. Our formulation is differentiable, except for the top-k gate selection, whose gradients we estimate with a rectangular straight-through estimator. This lets us learn which all-reduce operations to skip via backpropagation. After a brief joint fine-tuning of LoRA adapters and gates, our method nearly recovers the base model's performance. Skipping 50% of all-reduce operations, our method preserves 98.7% of the average downstream accuracy of Llama-3.1-70B (dense) and 98.2% for Qwen-2-57B (MoE).
est. 32% chance this paper gets accepted at ICLR 2027.
What do you think this paper will get?
All positions stay anonymous.