# How to Implement Distributed Training with MLX's Distributed Module: A Complete Guide

> Learn how to implement distributed training with MLX's distributed module. This guide covers tensor sharding and gradient synchronization for multi-device setups.

- Repository: [ml-explore/mlx](https://github.com/ml-explore/mlx)
- Tags: how-to-guide
- Published: 2026-06-18

---

**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:

```python
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`](https://github.com/ml-explore/mlx/blob/main/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 result
- **`all_gather`** – Concatenates tensors from all ranks along a specified axis
- **`send`** / **`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`](https://github.com/ml-explore/mlx/blob/main/python/mlx/nn/layers/distributed.py) uses an internal `_shard` helper that leverages `mx.split` and `mx.concatenate` to construct the local view:

```python
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`](https://github.com/ml-explore/mlx/blob/main/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:

```python
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:

```python
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`](https://github.com/ml-explore/mlx/blob/main/distributed.py)) follows the same communication patterns as their full-precision counterparts:

```python
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:

```python
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`](https://github.com/ml-explore/mlx/blob/main/python/mlx/nn/layers/distributed.py) | High-level Python API including sharding helpers and layer implementations |
| [`mlx/distributed/distributed.cpp`](https://github.com/ml-explore/mlx/blob/main/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 a `Group` object representing all MPI ranks.
- **Shard models** using `shard_inplace` or construct sharded layers directly with `AllToShardedLinear` and `ShardedToAllLinear`.
- **Synchronize gradients** automatically through the `sum_gradients` wrapper, which calls `all_sum` during backpropagation.
- **Use quantized variants** (`QuantizedAllToShardedLinear`, etc.) for memory-efficient distributed training.
- **Access the source** in [`python/mlx/nn/layers/distributed.py`](https://github.com/ml-explore/mlx/blob/main/python/mlx/nn/layers/distributed.py) and [`mlx/distributed/distributed.cpp`](https://github.com/ml-explore/mlx/blob/main/mlx/distributed/distributed.cpp) to 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`](https://github.com/ml-explore/mlx/blob/main/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.