How to Implement Distributed Training with MLX's Distributed Module: A Complete Guide
MLX enables distributed training across multiple devices through an MPI-backed Group abstraction that supports tensor sharding and automatic gradient synchronization.
MLX is Apple's machine learning framework designed for efficient computation on Apple Silicon. The framework includes a lightweight distributed training stack built on top of MPI, allowing you to shard model parameters across processes while maintaining synchronized gradients during training.
Core Architecture: The Group Abstraction
The foundation of distributed training with MLX is the Group object, which represents a set of processes that can exchange tensors.
To initialize the distributed runtime, call mlx.core.distributed.init(). This function starts the MPI runtime (or a no-op stub if MPI is unavailable) and returns the global Group covering all ranks:
import mlx.core as mx
from mlx.core import distributed as dist
# Initialize MPI and get the global group
group = dist.init()
print(f"Rank {group.rank()} of {group.size()}")
The C++ implementation of this initialization lives in mlx/distributed/distributed.cpp and is exposed to Python via the mlx.core.distributed module. Once initialized, use Group.rank() to identify the current process and Group.size() to determine the total number of processes.
Communication Primitives
MLX provides both collective and point-to-point communication operations that work directly on mx.array tensors:
all_sum– Sums tensors across all ranks and distributes the resultall_gather– Concatenates tensors from all ranks along a specified axissend/recv– Point-to-point communication between specific ranks
These primitives are implemented in backend-specific files such as mlx/backend/cuda/distributed.cu and analogous CPU/Metal implementations.
Sharding Model Parameters
The mlx.nn.layers.distributed module provides high-level utilities for splitting model parameters across your distributed group.
In-Place Model Sharding with shard_inplace
The function shard_inplace walks a module's parameter tree and slices each weight along a chosen axis (typically the output dimension), keeping only the local slice on each rank. The implementation in python/mlx/nn/layers/distributed.py uses an internal _shard helper that leverages mx.split and mx.concatenate to construct the local view:
import mlx.nn as nn
from mlx.nn.layers import distributed as ddist
# Create a simple MLP
model = nn.Sequential(
nn.Linear(1024, 4096),
nn.relu,
nn.Linear(4096, 1024),
)
# Shard all linear layers across the group (splits output dimension)
ddist.shard_inplace(model, sharding="all-to-sharded", segments=1, group=group)
After calling shard_inplace, each rank holds only 1/N of the original output dimensions, where N is the group size.
All-to-Sharded Linear Layers
AllToShardedLinear is a ready-made layer that computes partial outputs and returns them sharded across the group. According to the source code in python/mlx/nn/layers/distributed.py (lines 93-138), this layer accepts full input tensors but produces sharded outputs where each rank holds only its slice:
from mlx.nn.layers.distributed import AllToShardedLinear
# Create a sharded linear layer
sharded_fc = AllToShardedLinear(
input_dims=1024,
output_dims=4096,
bias=True,
group=group
)
# Forward pass returns sharded output
x = mx.random.randn(32, 1024) # batch of 32
y = sharded_fc(x) # shape: (32, 4096/size)
The backward pass automatically aggregates gradients across ranks using the sum_gradients wrapper, which ensures that all_sum is called during backpropagation.
Sharded-to-All Linear Layers
ShardedToAllLinear performs the inverse operation. It computes a local matrix multiplication, then calls mx.distributed.all_sum to combine partial results into a full tensor that every rank receives:
from mlx.nn.layers.distributed import ShardedToAllLinear
# Build a layer that produces replicated output
layer = ShardedToAllLinear(
input_dims=1024,
output_dims=4096,
bias=True,
group=group
)
# Forward returns full result on every rank
out = layer(x) # shape: (batch, 4096)
This is useful when you need to transition from sharded representations back to full tensors between layers.
Quantized Distributed Training
For memory-efficient distributed training, MLX provides quantized variants: QuantizedAllToShardedLinear and QuantizedShardedToAllLinear. These layers freeze parameters, use mx.quantize for storage efficiency, and rely on mx.quantized_matmul for the forward computation. The implementation (lines 558-613 in distributed.py) follows the same communication patterns as their full-precision counterparts:
from mlx.nn.layers.distributed import QuantizedAllToShardedLinear
qlayer = QuantizedAllToShardedLinear(
input_dims=1024,
output_dims=4096,
bias=True,
group=group
)
y = qlayer(x) # Uses quantized_matmul with automatic gradient aggregation
Implementing the Training Loop
Once your model is sharded, the training loop requires no special modifications. The sum_gradients wrapper (applied automatically by shard_inplace) ensures that gradients are synchronized across the group via all_sum during the backward pass:
optimizer = mx.optimizer.SGD(model.parameters(), lr=0.01)
for epoch in range(10):
for xb, yb in dataloader: # Same data on every rank
pred = model(xb) # Forward uses sharded layers
loss = mx.mean((pred - yb) ** 2)
loss.backward() # Gradient sum happens automatically
optimizer.step()
optimizer.zero_grad()
Because shard_inplace wrapped the linear layers with sum_gradients, each rank ends up with the same aggregated gradients and updates only its local parameter slice.
Key Implementation Files
The distributed training stack spans multiple files in the MLX repository:
| File | Description |
|---|---|
python/mlx/nn/layers/distributed.py |
High-level Python API including sharding helpers and layer implementations |
mlx/distributed/distributed.cpp |
Core MPI-backed implementation of group management and collectives |
mlx/backend/cuda/distributed.cu (and CPU/Metal analogs) |
Backend-specific kernels for communication operations |
docs/src/python/distributed.rst |
Documentation for the mlx.core.distributed namespace |
Summary
- Initialize distributed training with
mlx.core.distributed.init()to obtain aGroupobject representing all MPI ranks. - Shard models using
shard_inplaceor construct sharded layers directly withAllToShardedLinearandShardedToAllLinear. - Synchronize gradients automatically through the
sum_gradientswrapper, which callsall_sumduring backpropagation. - Use quantized variants (
QuantizedAllToShardedLinear, etc.) for memory-efficient distributed training. - Access the source in
python/mlx/nn/layers/distributed.pyandmlx/distributed/distributed.cppto understand the MPI-backed implementation.
Frequently Asked Questions
How do I initialize distributed training in MLX?
Call mlx.core.distributed.init() at the start of your script. This function initializes the MPI runtime and returns a Group object. You can then use group.rank() and group.size() to identify your process and the total number of processes. If MPI is not available, the function returns a no-op stub that allows single-process execution.
What is the difference between AllToShardedLinear and ShardedToAllLinear?
AllToShardedLinear accepts a full input tensor but produces a sharded output where each rank holds only a slice of the results (computed via segments or group.size()). ShardedToAllLinear accepts sharded inputs, performs local computation, and uses mx.distributed.all_sum to aggregate results so every rank receives the full output tensor. Use All-to-Sharded when splitting large output layers, and Sharded-to-All when you need to collect distributed representations.
Can I use quantization with distributed training in MLX?
Yes. MLX provides QuantizedAllToShardedLinear and QuantizedShardedToAllLinear in python/mlx/nn/layers/distributed.py. These layers freeze parameters, compress them using mx.quantize, and perform forward passes with mx.quantized_matmul. They maintain the same communication patterns as the standard layers, ensuring gradients are properly aggregated across ranks.
How are gradients synchronized across ranks in MLX?
Gradient synchronization happens automatically through the sum_gradients wrapper. When you shard a model using shard_inplace, it wraps the forward functions so that during the backward pass, all_sum is called to aggregate gradients across the group. This ensures that every rank receives the same gradient updates before the optimizer step, allowing each rank to update only its local parameter slice while maintaining consistency with the global model.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →