<!-- llms-explorer concept facts · https://llms-explorer.com/tree/tensor-parallel-all-reduce-count-per-transformer/ · pack 2026-10-05 · ~2088 tokens -->

# Tensor-parallel all-reduce count per transformer layer in MLX shard_linear

> 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 ...

Parent: [Mac local LLMs: Clusters, RDMA, exo and ds4](https://llms-explorer.com/tree/mac-local-llms-clusters-rdma-exo-ds4/) · 1 facets · 31 facts · page: https://llms-explorer.com/tree/tensor-parallel-all-reduce-count-per-transformer/

## 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](https://raw.githubusercontent.com/ml-explore/mlx/main/python/mlx/nn/layers/distributed.py)
- `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](https://raw.githubusercontent.com/ml-explore/mlx/main/python/mlx/nn/layers/distributed.py)
- `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](https://raw.githubusercontent.com/ml-explore/mlx/main/python/mlx/nn/layers/distributed.py)
- `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](https://raw.githubusercontent.com/ml-explore/mlx/main/python/mlx/nn/layers/distributed.py)
- 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](https://raw.githubusercontent.com/ml-explore/mlx-lm/main/mlx_lm/models/llama.py)
- 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](https://raw.githubusercontent.com/ml-explore/mlx-lm/main/mlx_lm/models/deepseek_v3.py)
- Qwen3-MoE follows the same pattern with `switch_mlp` sharded in place and a block-level `all_sum`, with the router left replicated. — [source](https://raw.githubusercontent.com/ml-explore/mlx-lm/main/mlx_lm/models/qwen3_moe.py)
- 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: `asserted`
- 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](https://raw.githubusercontent.com/exo-explore/exo/main/src/exo/worker/engines/mlx/auto_parallel.py)
- 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](https://raw.githubusercontent.com/exo-explore/exo/main/src/exo/worker/engines/mlx/auto_parallel.py)
- 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](https://raw.githubusercontent.com/ml-explore/mlx-lm/main/mlx_lm/models/deepseek_v3.py)
- 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: `asserted`
- 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](https://raw.githubusercontent.com/ml-explore/mlx-lm/main/mlx_lm/models/deepseek_v3.py)
- 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: `asserted`
- 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: `asserted`
- 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: `asserted`
- 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: `asserted`
- 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](https://raw.githubusercontent.com/ml-explore/mlx/main/python/mlx/nn/layers/distributed.py)
- 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: `asserted`
- 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: `asserted`
- `ShardedToAllLinear` and its quantized variant call `mx.distributed.all_sum` in `__call__`, then add the bias. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/python/mlx/nn/layers/distributed.py)
- `sum_gradients` is an identity forward and an all_sum only in its vjp, so all-to-sharded linears have no inference collective. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/python/mlx/nn/layers/distributed.py)
- `shard_inplace` adds no collective; `shard_linear` returns a new layer class that does the communication. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/python/mlx/nn/layers/distributed.py)
- mlx-lm's Llama makes 2 all_sums per layer (`o_proj` and `down_proj`). — [source](https://raw.githubusercontent.com/ml-explore/mlx-lm/main/mlx_lm/models/llama.py)
- mlx-lm's DeepSeek V3 makes 2 all_sums per dense layer and 2 per MoE layer. — [source](https://raw.githubusercontent.com/ml-explore/mlx-lm/main/mlx_lm/models/deepseek_v3.py)
- mlx-lm's Qwen3-MoE shards `switch_mlp` in place and all_sums once at the MoE block. — [source](https://raw.githubusercontent.com/ml-explore/mlx-lm/main/mlx_lm/models/qwen3_moe.py)
- exo's V4 TP path makes 2 all_sums per layer through `_AllSumLinear` and `ShardedMoEV4`. — [source](https://raw.githubusercontent.com/exo-explore/exo/main/src/exo/worker/engines/mlx/auto_parallel.py)
- mlx-lm pipeline parallelism ends with one `all_gather` of the final hidden state, not one per layer. — [source](https://raw.githubusercontent.com/ml-explore/mlx-lm/main/mlx_lm/models/deepseek_v3.py)
- The count of TP all_sums per token for a 61-layer MLX MoE is 122. — source: `asserted`
- 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: `asserted`
- Estimating all_sums from the number of sharded projections overcounts for MoE layers because in-place sharding adds none. — source: `asserted`
