DiLoCo (Distributed Low-Communication) Training#
This tutorial guides you through configuring and running DiLoCo and Streaming DiLoCo training in MaxText across multi-slice TPU clusters, multi-datacenter pods, and low-bandwidth DCN/WAN networks.
See also
This page is task-oriented: copy a recipe, adjust it, launch. For why the algorithm works, how to size your DCN link, and what each knob does mathematically, see the DiLoCo Theory & Mathematics Reference.
1. Vanilla DiLoCo vs. Streaming DiLoCo#
MaxText supports two modes of distributed low-communication training:
inner steps ───────────────────────────▶
Vanilla DiLoCo H = 4 here, typically 36–500
compute ████ ████ ████ ████ ████ ████ ████ ████
DCN ██████
▲
└─ whole model at once; compute waits
Streaming DiLoCo P = 4 fragments, Δh = 1, so H_eff = 4
compute ████ ████ ████ ████ ████ ████ ████ ████
DCN ▄ ▄ ▄ ▄ ▄ ▄ ▄
f1 f2 f3 f0 f1 f2 f3
└─ one fragment per step: same bytes per cycle,
1/P the size per transfer, no barrier
Key Differences:#
Vanilla DiLoCo (
enable_streaming_diloco=false):How it works: Each computing island trains independently for \(H\) inner steps (e.g., \(H=100\)). At every \(H\)-th step, training pauses for a global collective all-reduce where the entire model’s pseudo-gradient (\(\Delta \theta = \theta_{\text{outer}} - \theta_{\text{inner}}\)) is averaged across all islands over DCN and updated using outer Nesterov momentum.
When to use: Simpler baseline, ideal when \(H\) is large (e.g. \(H \ge 500\)) and the periodic all-reduce pause represents a negligible fraction of total training time.
Streaming DiLoCo (
enable_streaming_diloco=true):How it works: The model parameters are partitioned into \(P\) fragments (typically \(P = N_{\text{layers}} + 1\)). By setting \(H = P\), exactly 1 fragment is synchronized on every single local inner step (\(\Delta h = 1\)).
When to use: Optimal for high-throughput scaling across lower-bandwidth DCN/WAN networks, as it eliminates bursty communication spikes and removes the periodic step-\(H\) idle barrier.
Requirements:
scan_layers=true,num_diloco_fragments >= 2, andnum_decoder_layersdivisible bynum_diloco_fragments - 1. These are validated at startup.
Before your first run: the global batch is split, not replicated
With dcn_diloco_parallelism=K, each island trains on \(GBS/K\) tokens per inner step — adding islands does not multiply your token throughput per step, and \(GBS\) must be divisible by \(K\) or startup fails.
2. Prerequisites#
MaxText Environment: Follow the installation guide to set up your environment (
maxtext[tpu]ormaxtext[cuda12]).Compute Resources: A Google Kubernetes Engine (GKE) cluster with TPU slices managed via XPK.
Storage: A Google Cloud Storage (GCS) bucket for logging and Orbax checkpoints (
gs://<GCS_BUCKET>).
3. Production Recipe 1: Vanilla DiLoCo Multi-Slice Pre-training#
In this recipe, we train a model (e.g., Qwen3-8B) across 2 TPU v5p-128 slices using Vanilla DiLoCo with periodic synchronization every \(H=100\) steps:
python3 -m maxtext.trainers.pre_train.train \
run_name="vanilla-dlco-8b-01" \
base_output_directory="gs://your-bucket/maxtext-logs" \
dataset_path="gs://your-bucket/maxtext-datasets" \
dataset_name='c4/en:3.0.1' \
model_name="qwen3-8b" \
tokenizer_type=huggingface \
tokenizer_path=maxtext/assets/tokenizers/qwen3-tokenizer \
per_device_batch_size=8 \
max_target_length=2048 \
enable_diloco=true \
enable_streaming_diloco=false \
dcn_diloco_parallelism=2 \
diloco_sync_period=100 \
diloco_outer_lr=0.7 \
diloco_outer_momentum=0.9 \
steps=1000 \
enable_checkpointing=true \
checkpoint_period=100
Configuration Breakdown:#
enable_diloco=true: Enables outer optimization and multi-slice Low-Communication training acrossdcn_diloco_parallelism=2slices. The number of islands isnum_diloco_replicas = ici_diloco_parallelism * dcn_diloco_parallelism; setdcn_diloco_parallelism=-1to infer it fromnum_slices.enable_streaming_diloco=false: Disables parameter fragmentation and performs full-model pseudo-gradient all-reduce.diloco_sync_period=100: Islands execute 100 local AdamW steps independently before pausing to sync. Vanilla DiLoCo has nonum_diloco_fragments/scan_layersrequirement.diloco_outer_lr=0.7anddiloco_outer_momentum=0.9: Outer Nesterov momentum parameters (optax.sgd(..., nesterov=True)). Note the MaxText defaults are0.3and0.9.
4. Production Recipe 2: Streaming DiLoCo Dense Pre-training (Qwen3-8B)#
In this recipe, we train Qwen3-8B with Streaming DiLoCo across 2 TPU v5p-128 slices with \(H=P=37\) (synchronizing 1 fragment every step) via the SPMD runner script:
CLUSTER="mlperf-v5p" \
ZONE="europe-west4-b" \
PROJECT="cloud-tpu-multipod-dev" \
DEVICE_TYPE="v5p-128" \
NUM_SLICES="2" \
RUNNAME="stream-dlco-8b-01" \
XPK_WORKLOAD="stream-dlco-01" \
BASE_OUTPUT_DIRECTORY="gs://your-bucket/maxtext-logs" \
DATASET_PATH="gs://your-bucket/maxtext-datasets" \
MODEL_NAME="qwen3-8b" \
STEPS="1000" \
CHECKPOINT_PERIOD="100" \
DILOCO_SYNC_PERIOD="37" \
DILOCO_NUM_FRAGMENTS="37" \
DILOCO_NUM_COMM_OVERLAP_STEPS="0" \
DILOCO_USE_SEQUENTIAL_LAYERS="false" \
DILOCO_OUTER_LR="0.7" \
DILOCO_OUTER_MOMENTUM="0.9" \
bash src/maxtext/trainers/diloco/scripts/run_spmd_streaming_diloco.sh
The script builds a Docker image from your local working tree, pushes it, and submits the workload via XPK. It passes dcn_diloco_parallelism=${NUM_SLICES}, uses the Grain/TFRecord C4 pipeline, and pins the Qwen3 tokenizer. Add RESERVATION="<your-reservation>" if your cluster requires one.
The script’s defaults are not this tutorial’s recommendations
Omit these and you get different training behavior than the text describes:
Variable |
Script default |
Recommended here |
|---|---|---|
|
|
|
|
|
|
Both are set explicitly in the command above. The overlap-steps default of 2 is the more consequential one: it delays fragment application without any throughput benefit in the current SPMD design.
\(H = P = 37\) matches Qwen3-8B’s 36 decoder layers plus one fragment for the non-scanned embeddings and head, giving \(\Delta h = 1\) — one fragment synchronized per step.
5. Production Recipe 3: Streaming DiLoCo MoE Pre-training (Qwen3-30B-A3B)#
For large Mixture-of-Experts (MoE) architectures, this recipe demonstrates Streaming DiLoCo pre-training with the OLMo Grain data pipeline across 2x v5p-128 TPU slices:
XPK_CLUSTER="mlperf-v5p" \
XPK_ZONE="europe-west4-b" \
XPK_PROJECT="cloud-tpu-multipod-dev" \
XPK_DEVICE_TYPE="v5p-128" \
XPK_NUM_SLICES="2" \
RUN_NAME="qw3-olmo-dlco-01" \
WORKLOAD_NAME="qw3-olmo-01" \
BASE_OUTPUT_DIRECTORY="gs://your-bucket/maxtext-logs" \
OLMO_GCS_BASE="gs://your-bucket/datasets" \
MODEL_NAME="qwen3-30b-a3b" \
ENABLE_STREAMING_DILOCO="true" \
DILOCO_SYNC_PERIOD="49" \
DILOCO_NUM_FRAGMENTS="49" \
DILOCO_OUTER_LR="0.7" \
DILOCO_OUTER_MOMENTUM="0.9" \
bash src/maxtext/trainers/diloco/scripts/run_olmo_qwen3_30b_streaming_diloco.sh
49 Fragments: 48 MoE transformer decoder layers + 1 embedding/head fragment (\(H=49, P=49\)).
Grain Pipeline:
dataset_type=olmo_grain, reading a pre-built index (OLMO_INDEX_PATH, default/tmp/olmo-data/olmo/indices/olmo_index_seq8192.json). The script mountsOLMO_GCS_BASEwith gcsfuse atOLMO_LOCAL_MOUNTand remaps dataset paths from GCS to the local mount.Derived batch and schedule:
PER_DEVICE_BATCH_SIZEdefaults toTARGET_GLOBAL_BATCH / (devices_per_slice * XPK_NUM_SLICES), andSTEPSdefaults to a full pass overTOTAL_INSTANCES. Override any of these explicitly for shorter runs.
Same outer-LR caveat applies
This script also defaults DILOCO_OUTER_LR to 0.1. The recipe above overrides it to 0.7 explicitly.
6. Checkpointing & Resumption#
Automatic Resumption#
To resume an interrupted DiLoCo pre-training run, submit the workload with the same RUNNAME and BASE_OUTPUT_DIRECTORY:
RUNNAME="stream-dlco-8b-01" \
XPK_WORKLOAD="dlco-resm-01" \
BASE_OUTPUT_DIRECTORY="gs://your-bucket/maxtext-logs" \
STEPS="2000" \
bash src/maxtext/trainers/diloco/scripts/run_spmd_streaming_diloco.sh
Use a new XPK_WORKLOAD (workload names must be unique) but the same RUNNAME and BASE_OUTPUT_DIRECTORY, since the checkpoint directory is derived from those. MaxText detects the existing Orbax checkpoint, recognizes it as a multi-replica DiLoCo checkpoint, and restores the per-replica inner optimizer moments together with the outer Nesterov momentum trace.
Bootstrapping from Single-Slice Weights#
To initialize a multi-slice DiLoCo run from standard pre-trained single-slice weights, specify LOAD_FULL_STATE_PATH:
RUNNAME="stream-dlco-8b-boot" \
BASE_OUTPUT_DIRECTORY="gs://your-bucket/maxtext-logs" \
LOAD_FULL_STATE_PATH="gs://your-bucket/checkpoints/base_model/0/items" \
bash src/maxtext/trainers/diloco/scripts/run_spmd_streaming_diloco.sh
MaxText recognizes the restored state as a single-replica (non-DiLoCo) checkpoint, broadcasts the weights and optimizer state across all \(K\) islands along the diloco axis, and initializes a fresh outer Nesterov momentum state. The MoE recipe’s script additionally accepts LOAD_PARAMETERS_PATH for loading parameters only.
7. Tuning Guidelines#
Hyperparameter |
Recommended Setting |
Description |
|---|---|---|
|
|
Sync period. Setting \(H = P\) ensures exactly 1 fragment is synchronized every local step (\(\Delta h = 1\)). In general \(\Delta h = \max(1, \text{round}(H/P))\) and the effective period is \(P \cdot \Delta h\), so prefer an \(H\) that is a multiple of \(P\). |
|
|
Partition count (1 for non-scanned embeddings/head + 1 per transformer decoder layer). Must be \(\ge 2\), and |
|
|
Layer-to-fragment assignment. |
|
|
Outer learning rate. Start from |
|
|
Nesterov momentum coefficient for the outer optimizer. |
|
|
Delay in inner steps before applying outer weights. In the current SPMD design, this does not enhance hardware efficiency but simulates the algorithmic behavior of delayed weight merging; it will provide non-blocking hardware overlap in future MPMD multi-threading. Coupled with \(\alpha\). |
|
|
Soft parameter blending factor in \([0, 1]\) (\(\theta_{\text{inner}} \leftarrow \alpha \theta_{\text{inner}} + (1 - \alpha) \theta_{\text{outer}}\)). |
Choosing \(P\) for Your Model#
\(P = N_{\text{layers}} + 1\) is the recommended starting point, but any \(P\) satisfying \(P \ge 2\) and \(N_{\text{layers}} \bmod (P - 1) = 0\) is valid. Larger \(P\) means smaller, more frequent transfers:
Model |
\(N_{\text{layers}}\) |
Valid \(P\) values |
Recommended \(P\) |
Layers per fragment |
|---|---|---|---|---|
Qwen3-8B |
36 |
2, 3, 4, 5, 7, 10, 13, 19, 37 |
37 |
1 |
Qwen3-30B-A3B |
48 |
2, 3, 4, 5, 7, 9, 13, 17, 25, 49 |
49 |
1 |
A 32-layer model |
32 |
2, 3, 5, 9, 17, 33 |
33 |
1 |
Read the table as \(P - 1 \in \text{divisors}(N_{\text{layers}})\). If you need a \(P\) that doesn’t divide evenly, startup fails with an explicit error rather than silently rebalancing.
Practical Tuning Heuristics#
Synchronize Every Step (\(H = P = N_{\text{layers}} + 1\)): Setting
diloco_sync_periodequal tonum_diloco_fragmentswith \(P = N_{\text{layers}} + 1\) (e.g., \(H=37, P=37\) for 36-layer models like Qwen3-8B, or \(H=49, P=49\) for 48-layer models like Qwen3-30B) ensures a steady, constant stream of background communications by syncing 1 fragment on every local step.Keep \(H\) a Multiple of \(P\): \(\Delta h\) is computed as \(\max(1, \text{round}(H/P))\), so a period that doesn’t divide cleanly is silently rounded. \(H=100, P=37\) yields \(\Delta h = 3\) and an effective period of \(111\) steps — 11% longer than requested. Either set \(H = P\), or pick \(H\) as an exact multiple of \(P\).
Outer Learning Rate Tuning & Inverse Scaling Rule: Outer LR should be tuned alongside the inner optimizer learning rate. As a core heuristic:
\[\text{Higher Inner Learning Rate} \implies \text{Lower Outer Learning Rate}\]\[\text{Lower Inner Learning Rate} \implies \text{Higher Outer Learning Rate}\]When using standard AdamW inner optimization, starting with
diloco_outer_lr: 0.7is a strong baseline.Overlapping Steps (\(V\)) and Alpha (\(\alpha\)) in SPMD vs. MPMD:
num_communication_overlapping_steps(\(V\)) andcommunication_overlapping_alpha(\(\alpha\)) are coupled in defining the asynchronous weight merging policy:Current SPMD Design: Because JAX SPMD compiles each step into a synchronous XLA graph, setting \(V > 0\) or \(\alpha > 0\) does not enhance hardware training efficiency or hide network latency. However, it allows researchers to accurately simulate the algorithmic convergence behavior of delayed weight merging and soft parameter blending. Setting \(V=0\) is standard for performance.
Future MPMD Multi-Threading Design: In upcoming MPMD architectures with independent background communication threads, \(V\) and \(\alpha\) will provide true, non-blocking hardware compute/communication overlap.
8. Monitoring Convergence#
Aggregate scalars (learning/loss, throughput, etc.) are reported from replica 0. DiLoCo additionally emits a per-island loss for every replica:
learning/loss_island_0
learning/loss_island_1
...
learning/loss_island_{K-1}
Use these to diagnose the failure mode specific to low-communication training: islands drifting apart between synchronizations. A healthy run shows the per-island losses tracking each other closely and re-converging at each sync. A widening spread — especially one that does not shrink after a sync — usually means \(H\) is too large or the outer learning rate is too high.
9. Emulating a Slow Network#
To measure DiLoCo’s benefit on a cluster whose DCN is not actually the bottleneck, MaxText can throttle per-VM egress using a Linux traffic-control token bucket filter:
Config |
Default |
Description |
|---|---|---|
|
|
Per-VM egress limit, e.g. |
|
|
Token bucket filter burst size. |
|
|
Token bucket filter latency threshold. |
|
|
Network interface the shaping rules are applied to. |
This lets you compare synchronous data parallelism against DiLoCo at a realistic WAN bandwidth before committing to a cross-datacenter run.
10. Troubleshooting#
All DiLoCo misconfigurations are caught at startup, before any accelerator work begins. The table maps each error to its cause.
Error message (abridged) |
Cause and fix |
|---|---|
|
Streaming is a mode of DiLoCo, not a replacement. Set both flags. |
|
|
|
You need at least one fragment for non-scanned params and one for layers. |
|
Fragments are slices of the scanned layer stack; unscanned models have no such stack. Enable |
|
Pick \(P\) such that \(P - 1\) divides \(N_{\text{layers}}\). See the table in §7. |
|
The global batch is split across islands. Adjust |
|
|
Run diverges or loss spikes at sync boundaries. Lower diloco_outer_lr first — it is the most sensitive knob, and the inverse scaling rule above means a high inner LR needs a lower outer LR. If per-island losses spread apart steadily, reduce diloco_sync_period so islands reconcile more often.
No throughput gain versus synchronous training. Confirm the DCN link was actually the bottleneck. On a fast interconnect, DiLoCo mainly adds outer-state memory without a speedup; use dcn_bandwidth_limit (§9) to verify the benefit under realistic bandwidth before committing to a multi-datacenter run.
11. Next Steps#
Deep dive into mathematical foundations: DiLoCo Theory & Mathematics Reference.
Explore input pipeline options: Data Input Pipeline Guides.
Learn about sharding strategies on TPUs: Sharding on TPUs.