Tensor-parallel all-reduce count per transformer layer in MLX shard_linear
Parent: Mac local LLMs: Clusters, RDMA, exo and ds4 · Published reference · snapshot 2026-10-05
↓ Facts as markdownall context files
A tensor-parallel all-reduce in MLX is an `mx.distributed.all_sum` on an activation. The count per transformer layer equals the number of `ShardedToAllLinear` (or quantized equivalent) calls made by layers sharded with `shard_linear`, plus the explicit `all_sum` calls that model code adds around ...
These notes link each claim to its source. A source may be a research report hosted on this site rather than the primary document. A published reference means the content is available; it does not certify independent review or accuracy.Read the editorial policy and follow the sources before relying on a claim.
Facts
- A tensor-parallel all-reduce in MLX is an `mx.distributed.all_sum` on an activation. The count per transformer layer equals the number of `ShardedToAllLinear` (or quantized equivalent) calls made by layers sharded with `shard_linear`, plus the explicit `all_sum` calls that model code adds around a block that was sharded in place. [source]
- `ShardedToAllLinear.__call__` multiplies by the local weight slice, calls `mx.distributed.all_sum(x, group)`, then adds the bias once after the sum. The quantized class does the same around `mx.quantized_matmul`. So each "sharded-to-all" linear costs exactly one all_sum per call. [source]
- `AllToShardedLinear.__call__` applies `sum_gradients` to its input, and `sum_gradients` is an identity in the forward pass with the all_sum only in its vjp. An "all-to-sharded" linear therefore costs no collective in inference. [source]
- `shard_inplace` only slices the weights of an existing module and adds no collective. `shard_linear` returns a new layer class that carries the all_sum (forward for sharded-to-all, backward for all-to-sharded). [source]
- Dense Llama in mlx-lm shards q, k, v, gate and up as all-to-sharded and `o_proj` and `down_proj` as sharded-to-all with `shard_linear`, so it makes 2 all_sums per layer and none at the embedding or LM head. [source]
- DeepSeek V3 in mlx-lm: `o_proj` is sharded-to-all (1 all_sum); a dense MLP layer uses `shard_linear` for gate, down and up (1 more, 2 total); a MoE layer shards `shared_experts` and `switch_mlp` in place and the MoE block calls `mx.distributed.all_sum` on its output once (1 more, 2 total). The code comment reads "Shard in place since the MoE should be responsible for aggregating the results". [source]
- Qwen3-MoE follows the same pattern with `switch_mlp` sharded in place and a block-level `all_sum`, with the router left replicated. [source]
- Sharding the shared expert in place folds its output into the routed experts' sum, so a MoE layer with routed and shared experts needs one MoE all_sum, not two. [source]
- exo's V3, V3.2 and Kimi K2.5 strategy uses the same wrappers; its V4 path makes the attention all_sum with `_AllSumLinear` (after the in-place `wo_a`) and the MoE all_sum with `ShardedMoEV4`, so it also makes 2 per layer. [source]
- exo's V4 docstring rejects a design with "61 extra all_gathers/token" in favor of replicating `wo_b`, which also confirms 61 transformer layers for that model. [source]
- Pipeline mode has a different collective profile: `send` and `recv` at stage boundaries, plus one `all_gather` of the final hidden state per forward pass (not per layer). [source]
- Per decoded token a 61-layer model therefore makes 122 all_sums, each carrying an activation of batch x hidden x dtype bytes (14 KB for hidden size 7168 in bf16 at batch 1). [source]
- The all_sum placement follows a design note in the DeepSeek V3 shard code: MoE layers shard "in place since the MoE should be responsible for aggregating the results"; the note gives no further reason or date. [source]
- Counting calls to `shard_linear` is not enough: the same projection costs one all_sum if sharded with `shard_linear` and none if sharded with `shard_inplace`. Any estimate from "number of sharded matrices" overcounts for MoE layers. [source]
- A model that shards the attention output with `shard_linear` and the MLP with `shard_inplace` plus a block wrapper still lands at 2 per layer. A model whose attention output projection is replicated (as `wo_b` in exo's V4 path) can still cost one if an `_AllSumLinear` wraps it. [source]
- If a collective sits on the `o_proj` and the MLP is split further (for example separate shared and routed all_sums) the count rises to 3 per layer; no mlx-lm model read here does that. [source]
- The mlx-benchmarks modeled TP2 overhead (about 36% of token time over Thunderbolt RDMA for Llama 405B) assumes one all-reduce per layer, so with two per layer the modeled communication share roughly doubles. [source]
- One versus two per layer. mlx-benchmarks INTERCONNECTS says 1; the MLX primitives and mlx-lm model files give 2 for dense Llama and for both dense and MoE DeepSeek V3 layers. The primitives are the primary source. [source]
- Whether MLX will fuse the attention and MLP all_sums, as some GPU stacks do with a parallel-block layout; no PR was seen. [source]
- Whether `mx.distributed.all_sum` on an 8-16 KB tensor over JACCL has a fixed per-call floor that dominates the count-times-bytes estimate (see per-collective-all-sum-latency-in-decode-over-ja.md). [source]
- `ShardedToAllLinear` and its quantized variant call `mx.distributed.all_sum` in `__call__`, then add the bias. [source]
- `sum_gradients` is an identity forward and an all_sum only in its vjp, so all-to-sharded linears have no inference collective. [source]
- `shard_inplace` adds no collective; `shard_linear` returns a new layer class that does the communication. [source]
- mlx-lm's Llama makes 2 all_sums per layer (`o_proj` and `down_proj`). [source]
- mlx-lm's DeepSeek V3 makes 2 all_sums per dense layer and 2 per MoE layer. [source]
- mlx-lm's Qwen3-MoE shards `switch_mlp` in place and all_sums once at the MoE block. [source]
- exo's V4 TP path makes 2 all_sums per layer through `_AllSumLinear` and `ShardedMoEV4`. [source]
- mlx-lm pipeline parallelism ends with one `all_gather` of the final hidden state, not one per layer. [source]
- The count of TP all_sums per token for a 61-layer MLX MoE is 122. [source]
- The one-per-layer figure in mlx-benchmarks undercounts, and its modeled TB5 communication share is about half of what two per layer would give. [source]
- Estimating all_sums from the number of sharded projections overcounts for MoE layers because in-place sharding adds none. [source]
Children
- No children recorded.