Sharding technique speeds up transformer retraining and distillation
Affinity-Aware Sharding for Delayed Tensor Parallelism
Machine LearningComputation and LanguageDistributed, Parallel, and Cluster Computing
Summary
Transformer models are large AI programs that can be split across many devices to work faster. One way to do this, called Delayed Tensor Parallelism (DTP), changes how parts of the model communicate, which requires retraining the model after splitting it differently. The authors found that rearranging the parts shared on each device to match how they interact speeds up this retraining. They show a method to find the best arrangement of these parts, making the retraining process about twice as fast on tested models.
What this means in practice
- •For deep learning engineers: Speed up retraining and distillation when adapting transformer models to delayed tensor parallelism by arranging neuron and attention head assignments optimally.
- •For infrastructure teams in ai companies: Reduce GPU time and resource use for transformer model updates by applying affinity-aware sharding during model partitioning.
Authors
Eloi de Reynal
Abstract
Delayed Tensor Parallelism (DTP) removes the blocking all-reduce of tensor-parallel Transformer inference. Every device adds its own partial output to its residual stream (and broadcasts it) immediately, but only gathers (receives) the other devices' partials $δ$ modules later. A TP to DTP change therefore amounts to a real architecture change, and dense Transformer models need to be retrained or distilled after adaptation. We show that DTP breaks the permutation symmetry of neurons inside FFNs and of KV heads inside attention modules, and that this symmetry breakage makes the sharding itself a modelling decision. We show that maximising the affinity between the KV heads and the FFN neurons co-located on a device, by permuting the dense model before sharding, speeds up the distillation or retraining process. The affinity is measured with a first-order approximation of the damage that losing a head's contribution does to each neuron's output, and the co-located affinity is maximised with a coordinate-ascent optimiser that alternates an exact balanced assignment of neurons with an exhaustive search over the KV head partitions. The whole procedure takes under two minutes on one GPU for Qwen3-0.6B and Danube3-500M. On these models, at $δ=1$, the affinity-optimised layouts reach any distillation target in about half to two thirds of the steps needed by the naive contiguous layouts, over the whole 10k-step range we tested, and every optimised seed beats every contiguous seed and all but one of the sixteen random layouts. We also show that the co-located affinity score at initialisation predicts the KL to the base model after training, across seventeen layouts ranging from anti-optimised to optimised (Pearson $-0.81$ and $-0.89$).