In Part 1, I covered the batch size hierarchy: the four levels from micro-batch to global tokens, auto-computed gradient accumulation, and why the progress bar lies to you. That was the arithmetic of distributed training. It tells you how much each GPU processes per step.
This post is about the harder question: how do 256 GPUs coordinate?
I’ll cover three things. First, DDP and FSDP: the two strategies for distributing model state across GPUs, and when you’re forced to switch from one to the other. Second, data sharding: why you can’t just let every GPU read from the same filesystem, and the surprisingly tricky epoch calculation bug that follows. Third, the multi-node launch pattern: the two-line mistake that spawns 2,048 processes instead of 256.
DDP: the simple case
DDP stands for Distributed Data Parallel. It’s the default strategy, the one most tutorials teach, and the one you should use until it stops working. The idea is straightforward.
Every GPU holds a complete copy of the model. All the parameters, all the optimizer states, all the gradients. Every GPU processes a different mini-batch of data, computes gradients independently, and then: the only moment of communication. All GPUs run an allreduce to average their gradients. After that, each GPU applies the same averaged gradient to its own copy of the weights. Since all copies started identical and all received the same averaged gradient, they remain identical.
One communication per optimizer step. That’s it. The rest is embarrassingly parallel.
The memory budget
Here’s what one GPU needs to hold in DDP for a 1B parameter model:
| Component | Size | Why |
|---|---|---|
| Parameters (bf16) | ~2 GB | 1B params × 2 bytes |
| Gradients (bf16) | ~2 GB | Same shape as params |
| Optimizer states (fp32) | ~8 GB | AdamW: 2 states per param × 4 bytes |
| Activations | ~15-20 GB | Depends on batch size and seq length |
| Total | ~27-32 GB |
Each compute die on our cluster has 64 GB of HBM. A 1B model fits comfortably. Even with per_device_batch_size = 8 and 1024-length sequences, we’re using about half the available memory. DDP is the right choice here: minimal communication overhead, maximum simplicity.
But what happens when the model doesn’t fit?
FSDP: when DDP runs out of memory
An 8B parameter model needs roughly 16 GB just for parameters in bf16. Gradients: another 16 GB. AdamW optimizer states in fp32: 64 GB. We’re already at 96 GB before a single activation is stored. That’s 1.5x more than the entire GPU memory. DDP is physically impossible.
FSDP (Fully Sharded Data Parallel) solves this by sharding everything. Parameters, gradients, and optimizer states are split across all GPUs. Each GPU holds only $\frac{1}{W}$ of the model, where $W$ is the world size.
For an 8B model on 256 GPUs:
| Component | Per-GPU (DDP) | Per-GPU (FSDP) | Reduction |
|---|---|---|---|
| Parameters | 16 GB | 0.06 GB | 256× |
| Gradients | 16 GB | 0.06 GB | 256× |
| Optimizer states | 64 GB | 0.25 GB | 256× |
| Total (excl. activations) | 96 GB | 0.37 GB |
From 96 GB to 0.37 GB. The 8B model now fits on a single 64 GB compute die with plenty of room for activations.
The catch is communication. DDP communicates once per optimizer step (allreduce the gradients). FSDP communicates constantly:
- Before each layer’s forward pass: allgather the full parameters for that layer from all GPUs (because each GPU only has $\frac{1}{W}$ of them)
- After each layer’s forward pass: discard the gathered parameters (free the memory for the next layer)
- Before each layer’s backward pass: allgather the parameters again (need them for gradient computation)
- After each layer’s backward pass: reduce-scatter the gradients (each GPU keeps $\frac{1}{W}$ of the averaged gradient)
That’s two allgathers and one reduce-scatter per layer. For a 32-layer model, that’s 96 collective operations per optimizer step, compared to DDP’s one. The communication volume is also higher: FSDP moves full parameter tensors, not just gradients.
When to switch
The decision is simple in practice. Try DDP first. If you get an out-of-memory error, switch to FSDP. For our models:
| Model | Strategy | Why |
|---|---|---|
| 240M | DDP | Fits easily, ~6 GB total |
| 500M | DDP | Fits fine, ~12 GB total |
| 1B | DDP | Still fits, ~30 GB total |
| 8B | FSDP | 96 GB without activations: doesn’t fit |
There’s a middle ground too. FSDP lets you shard only optimizer states and gradients while keeping full parameters on each GPU (this is what the older ZeRO Stage 1 and 2 did). But in my experience, once you need FSDP at all, you might as well go full sharding. The communication overhead difference between partial and full sharding is smaller than you’d expect, and full sharding gives you the most memory headroom for larger batch sizes.
Activation checkpointing: the other memory lever
Even with FSDP, activations can blow up memory. During the forward pass, every layer saves its activations for the backward pass. For a 32-layer model with batch size 8 and sequence length 1024, that’s a lot of tensors sitting in memory simultaneously.
Activation checkpointing trades compute for memory: instead of saving all activations, save only the inputs at each layer boundary. During the backward pass, recompute the forward activations on the fly. You lose about 33% in throughput (roughly one extra forward pass worth of compute), but you reduce activation memory from $O(N)$ to $O(\sqrt{N})$: roughly a 60-80% reduction depending on model depth.
For the 8B model, activation checkpointing is non-negotiable. Even with FSDP sharding, the activations alone would exceed 64 GB at reasonable batch sizes. The FSDP config enables it at the framework level:
fsdp_config:
fsdp_activation_checkpointing: true
fsdp_auto_wrap_policy: TRANSFORMER_BASED_WRAP
fsdp_reshard_after_forward: true
fsdp_reshard_after_forward: true is the key line that tells FSDP to free gathered parameters immediately after each layer’s forward pass. Without it, FSDP would keep all gathered parameters in memory: defeating the entire purpose.
One gotcha that cost me an hour: you cannot set both fsdp_activation_checkpointing: true in the accelerate config and gradient_checkpointing = True in the training arguments. They’re the same thing exposed through two different APIs, and setting both causes a cryptic error. Pick one. If you’re using FSDP, use the FSDP config’s version.
The allreduce: ring, tree, and why it matters
I glossed over the allreduce earlier. It deserves a closer look, because at 256 GPUs across 32 nodes, the choice of allreduce algorithm meaningfully affects training speed.
An allreduce takes $N$ tensors (one per GPU) and produces the average on all GPUs. Two common algorithms:
Ring allreduce: arrange GPUs in a logical ring. Each GPU sends a chunk to its neighbor, receives a chunk from the other side, accumulates. After $2(N-1)$ steps, every GPU has the full averaged result. Communication volume is optimal: each GPU sends and receives exactly $\frac{2(N-1)}{N} \times D$ bytes, where $D$ is the tensor size. But latency is $O(N)$: the ring has to go around twice.
Tree allreduce: arrange GPUs in a binary tree. Reduce up the tree (log N steps), broadcast down (log N steps). Total latency is $O(\log N)$, but each link carries the full tensor, so bandwidth utilization is lower.
At 256 GPUs: ring needs 510 steps, tree needs about 16. We use tree. The latency wins are enormous at this scale, and the high-speed Slingshot interconnect (4 NICs per node, 100 GB/s aggregate) has enough bandwidth that the tree’s lower bandwidth efficiency isn’t the bottleneck.
# From our RCCL tuning config:
export NCCL_ALGO=Tree
export NCCL_CROSS_NIC=1 # use all 4 NICs
export NCCL_NET_GDR_LEVEL=3 # GPU-Direct RDMA: GPU→NIC, bypass CPU
export NCCL_MIN_NCHANNELS=32 # more parallel channels
GPU-Direct RDMA is worth calling out: it lets one GPU write directly to another GPU’s memory across the network, without staging through CPU memory. On our cluster, this means a gradient tensor on node 0’s GPU can be sent to node 31’s GPU without either node’s CPU touching the data. The CPU is completely out of the critical path.
Data sharding: the filesystem problem
Here’s a problem I didn’t anticipate. Our training dataset is about 200 GB: 97.6 billion tokens of pre-tokenized text, stored on the parallel filesystem. When I first ran a 32-node job, all 256 GPUs tried to read from the same filesystem simultaneously.
The job crashed within minutes.
Not an out-of-memory crash. A filesystem crash. The parallel filesystem couldn’t handle 256 concurrent readers hammering it with random-access reads. I/O contention caused node-level timeouts, which the communication library interpreted as node failures, which triggered a cascade of process exits.
The fix is data sharding. Instead of everyone reading from the shared filesystem, we pre-split the dataset into 128 shards and copy each node’s shards to its local NVMe SSD before training starts. Each node reads exclusively from its own local storage. Zero filesystem contention.
Shared filesystem (200 GB):
shard_000.bin ... shard_127.bin (128 shards)
Node 0's local NVMe:
shard_000.bin, shard_001.bin, shard_002.bin, shard_003.bin
Node 1's local NVMe:
shard_004.bin, shard_005.bin, shard_006.bin, shard_007.bin
...
Node 31's local NVMe:
shard_124.bin, shard_125.bin, shard_126.bin, shard_127.bin
128 shards across 32 nodes gives 4 shards per node. Each shard is about 762 million tokens. Each node’s 8 GPUs share those 4 local shards, with a DistributedSampler ensuring no two GPUs on the same node process the same sample.
This is a clean solution. But before we get to the subtle bug it introduces, let me talk about the tool that made this whole data pipeline survivable.
tmux: the tool nobody tells you about early enough
Downloading 200 GB of data to a cluster takes hours. Tokenizing it takes longer. Copying shards to local NVMe across 32 nodes: more hours. And your SSH connection will drop at some point. It always does. Network hiccup, laptop goes to sleep, VPN reconnects. The download dies halfway through, and you start over.
I lost an entire day to this before Benjamin Therien at Mila showed me tmux. Ben is one of the strongest engineers I’ve worked with: the kind of person who makes infrastructure problems disappear before you even notice them. It was one of those moments where you wonder how you ever worked without it.
tmux runs a persistent terminal session on the server. You start it, launch your download, detach from it, close your laptop, go get coffee. Come back, reattach. Everything is still running. Your SSH connection is just a window into the session: the session itself lives on the server and doesn’t care if you disconnect.
Here’s the workflow that saved me:
# Start a named session
tmux new -s data-download
# Inside tmux: start your long-running job
python tokenize_dataset.py --output /scratch/shards/
# Detach: Ctrl+b, then d
# Your laptop can now disconnect. The job keeps running.
# Later, reattach from any SSH session:
tmux attach -t data-download
# Everything is exactly where you left it.
Once I had tmux running, I could kick off a 6-hour download before bed, close my laptop, and check the progress in the morning. No restarts. No partial downloads. No wasted hours.
Beyond data downloads, tmux became my default way of working on clusters. I keep a session for training runs, another for monitoring, another for debugging. Each one survives disconnects. It completely changed how I interact with remote machines.
Some commands worth knowing:
# Session management
tmux new -s name # new named session
tmux ls # list sessions
tmux attach -t name # reattach
tmux kill-session -t name # clean up
# Inside tmux (all start with Ctrl+b)
Ctrl+b d # detach
Ctrl+b c # new window
Ctrl+b n / p # next / previous window
Ctrl+b % # split pane vertically
Ctrl+b " # split pane horizontally
Ctrl+b arrow # switch between panes
Ctrl+b [ # scroll mode (q to exit)
The split panes are especially useful: training logs in one pane, GPU usage in another, a shell for debugging in a third. All in one tmux session, all persistent.
I genuinely think tmux is one of those tools that should be taught on day one of any HPC onboarding. It’s that fundamental. Thanks, Ben.
The epoch calculation bug
When you shard data across nodes, len(dataset) on any given node returns the local sample count, not the global count. My node has 4 out of 128 shards. It sees about 6.1 million samples. The full dataset has 195 million.
The training script computes max_steps from the dataset size:
one_epoch_steps = len(dataset["train"]) // global_batch_samples
With node-sharded data, this gives:
\[\frac{6{,}103{,}512}{4096} = 1{,}490 \text{ steps}\]1,490 steps. That’s $\frac{1}{32}$ of the correct value. The script thinks one epoch is 1,490 steps and caps training there. Instead of seeing 100 billion tokens, the model sees 3 billion and stops. Everything looks fine in the logs: “training complete, 1 epoch finished.” But the model is wildly undertrained.
I caught this because the loss was suspiciously high at the end of “training.” It hadn’t converged at all. Took me a while to trace it back to the epoch calculation.
The fix is one line:
if node_sharded:
global_train_samples = len(dataset["train"]) * num_nodes
else:
global_train_samples = len(dataset["train"])
one_epoch_steps = global_train_samples // global_batch_samples
# = 195,312,384 // 4096 = 47,683 steps ← correct
Multiply the local count by the number of nodes. That’s it. But if you don’t know to look for it, this bug is invisible. The training completes successfully, the loss decreases, checkpoints are saved. Everything appears normal. The only clue is that training is 32x shorter than it should be.
Launching 256 processes without launching 2,048
Multi-node distributed training on a Slurm cluster involves two layers of process management: Slurm launches processes across nodes, and then a framework like Accelerate (or torchrun) spawns GPU workers within each node. Getting the interaction right between these two layers is where everyone makes the same mistake at least once.
The wrong way:
srun -n 256 accelerate launch --num_processes 8 train.py
This tells Slurm to launch 256 tasks. Each task runs accelerate launch, which spawns 8 GPU workers. Total processes: $256 \times 8 = 2{,}048$. All 2,048 try to bind to port 29500 for the distributed backend. 2,040 of them fail with “address already in use.” The job crashes immediately.
The right way:
srun --ntasks-per-node=1 --gpus-per-node=8 \
bash -c "accelerate launch \
--machine_rank \$SLURM_NODEID \
--main_process_ip ${MASTER_ADDR} \
--num_machines 32 \
--num_processes 256 \
train.py"
Slurm launches 1 task per node (32 total). Each task runs accelerate launch, which spawns 8 GPU workers on that node. Total: $32 \times 8 = 256$. Exactly right.
The key insight is that Slurm and Accelerate operate at different levels. Slurm handles inter-node distribution. Accelerate handles intra-node distribution. If you let both try to do everything, they multiply instead of composing.
Rendezvous: how 256 processes find each other
All 256 processes need to establish communication before training can begin. This is the rendezvous. One node is designated as the master, and all processes connect to it:
MASTER_ADDR=$(getent ahostsv4 "$(scontrol show hostname $SLURM_NODELIST \
| head -n 1)" | head -n 1 | awk '{print $1}')
MASTER_PORT=3442
Two details here that took me embarrassingly long to learn.
First: the master address must be IPv4. On our cluster, hostname resolution sometimes returns an IPv6 address first. The communication library doesn’t support IPv6 and fails with errno:97: Address family not supported by protocol. The error message is completely unhelpful. I spent two hours reading RCCL source code before I thought to check getent ahosts and saw the IPv6 address. Forcing IPv4 resolution with getent ahostsv4 fixed it instantly.
Second: the master port must be unused. This sounds obvious, but on a shared cluster, another user’s job might have recently released a port that’s still in TIME_WAIT state. We hardcode port 3442 (an uncommon choice) and validate it’s available before training starts.
The full communication stack
Let me put the whole stack together, from hardware to framework:
Application layer: PyTorch DDP / FSDP
↓
Communication layer: RCCL (AMD's collective communication library)
↓
Transport layer: libfabric (Slingshot CXI provider)
↓
Hardware: 4× Slingshot-11 NICs per node (25 GB/s each)
GPU-Direct RDMA (GPU memory ↔ NIC, bypass CPU)
The tuning knobs that matter most at our scale:
NCCL_SOCKET_IFNAME=hsn # bind to Slingshot NICs, not loopback
NCCL_CROSS_NIC=1 # aggregate bandwidth across all 4 NICs
NCCL_NET_GDR_LEVEL=3 # GPU-Direct RDMA level
NCCL_ALGO=Tree # tree allreduce for 32+ nodes
NCCL_MIN_NCHANNELS=32 # parallel communication channels
NCCL_SOCKET_FAMILY=AF_INET # force IPv4
FI_CXI_ATS=0 # disable address translation overhead
Each of these was a separate debugging session. The default values work fine for single-node training. They silently degrade at scale.
Launching 256 processes without launching 2,048 (visual)
Tracing one optimizer step through the full stack
Let me close by tracing one complete optimizer step for a 1B model on 32 nodes with DDP. Everything from Parts 1 and 2, end to end.
Step #10000 begins
├── Gradient accumulation loop (2 iterations):
│
│ Iteration 1:
│ ├── Each node reads a micro-batch from local NVMe (no filesystem contention)
│ ├── DistributedSampler ensures 8 GPUs on each node get different samples
│ ├── GPU 0: forward(8 samples) → loss → backward → local grad
│ ├── GPU 1: forward(8 samples) → loss → backward → local grad
│ ├── ...
│ ├── GPU 255: forward(8 samples) → loss → backward → local grad
│ └── 2048 samples processed. No communication yet.
│
│ Iteration 2:
│ ├── Same as above, 2048 more samples
│ ├── Gradients accumulated locally: grad += new_grad
│ └── 4096 samples total across all GPUs
│
├── AllReduce (the ONLY communication):
│ ├── RCCL tree allreduce across 256 GPUs
│ ├── Gradient tensors flow through Slingshot via GPU-Direct RDMA
│ ├── Tree depth: ~8 levels (log₂ 256)
│ ├── Each GPU now holds the averaged gradient
│ └── ~50ms for a 1B model's gradients
│
├── Optimizer step:
│ ├── Each GPU runs AdamW independently
│ ├── Identical inputs (same avg gradient, same state) → identical outputs
│ └── All 256 model copies remain synchronized
│
└── Logging:
[step 10000 epoch 0.21] loss=2.54 tflops=27.3 mfu=28.5%
tokens_seen: 41.9B
4,194,304 tokens digested. 4096 samples from 256 GPUs across 32 nodes, read from local NVMe, processed independently, gradients averaged once over Slingshot, weights updated identically everywhere. About 3.5 seconds.
Then it does it again. 23,841 more times.
References
If you want to go deeper on any of these topics, here’s what I found most useful.
DDP and FSDP:
- Li et al., PyTorch Distributed: Experiences on Accelerating Data Parallel Training (2020). The original DDP paper. Section 3 on gradient bucketing and communication overlap is the part that matters.
- Zhao et al., PyTorch FSDP: Experiences on Scaling Fully Sharded Data Parallel (2023). The FSDP paper. Read Section 4 on the memory/communication trade-off.
- Rajbhandari et al., ZeRO: Memory Optimizations Toward Training Trillion Parameter Models (2019). The DeepSpeed ZeRO paper that FSDP builds on. The three stages of sharding (optimizer states → gradients → parameters) map directly to FSDP’s configuration options.
Allreduce algorithms:
- Thakur et al., Optimization of Collective Communication Operations in MPICH (2005). Still the best reference for ring vs tree vs recursive halving/doubling. Old paper, but the algorithms haven’t changed.
- The NCCL documentation on collective operations. Written for NVIDIA’s library, but RCCL (AMD’s fork) implements the same algorithms and environment variables.
Scaling laws and token budgets:
- Hoffmann et al., Training Compute-Optimal Large Language Models (2022). The Chinchilla paper. The key result is $D \approx 20N$: train on 20 tokens per parameter.
- Muennighoff et al., Scaling Data-Constrained Language Models (2023). Why training past Chinchilla-optimal (like our 5× over) still helps for smaller models.
Activation checkpointing:
- Chen et al., Training Deep Nets with Sublinear Memory Cost (2016). The original gradient/activation checkpointing paper. The $\sqrt{N}$ memory bound is elegant.
Practical guides:
- The HuggingFace Accelerate documentation for the DDP/FSDP configuration layer.
- The PyTorch FSDP tutorial for worked examples of the wrapping policy and sharding strategy.