Affinity-Aware Sharding for Delayed Tensor Parallelism
Mirrored from arXiv — Machine Learning for archival readability. Support the source by reading on the original site.
Computer Science > Machine Learning
Title:Affinity-Aware Sharding for Delayed Tensor Parallelism
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 $\delta$ 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 $\delta=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$).
| Comments: | 16 pages, 11 figures, 3 tables |
| Subjects: | Machine Learning (cs.LG); Computation and Language (cs.CL); Distributed, Parallel, and Cluster Computing (cs.DC) |
| ACM classes: | I.2.7; C.1.4 |
| Cite as: | arXiv:2609.13846 [cs.LG] |
| (or arXiv:2609.13846v1 [cs.LG] for this version) | |
| https://doi.org/10.48550/arXiv.2609.13846
arXiv-issued DOI via DataCite (pending registration)
|
Access Paper:
- View PDF
- HTML (experimental)
- TeX Source
Current browse context:
References & Citations
Bibliographic and Citation Tools
Code, Data and Media Associated with this Article
Demos
Recommenders and Search Tools
arXivLabs: experimental projects with community collaborators
arXivLabs is a framework that allows collaborators to develop and share new arXiv features directly on our website.
Both individuals and organizations that work with arXivLabs have embraced and accepted our values of openness, community, excellence, and user data privacy. arXiv is committed to these values and only works with partners that adhere to them.
Have an idea for a project that will add value for arXiv's community? Learn more about arXivLabs.
More from arXiv — Machine Learning
-
Stable and Faithful Explanations for Knowledge Tracing
Sep 25
-
SMILESGNN: Interpretable Clinical Toxicity Prediction via SMILES-Graph Cross-Attention Fusion
Sep 25
-
CFD Correction of Open Tip Clearance Flow in a Compressor Cascade Using VAE Latent Space Adaptation
Sep 25
-
CARE: Condition-Aware Representation Regularization for Diffusion Models
Sep 25
Discussion (0)
Sign in to join the discussion. Free account, 30 seconds — email code or GitHub.
Sign in →No comments yet. Sign in and be the first to say something.