Parallelism strategies for distributed training#
This guide explains the parallelism dimensions used when training large foundation models on AMD GPUs with Primus. It moves from basic data parallelism to advanced combinations of tensor, pipeline, context, and expert parallelism, including how Primus exposes these options through Megatron-LM and TorchTitan.
For Megatron YAML flags and environment tuning, see Megatron parameters and Environment variables.
1. Introduction#
Why parallelism is needed#
Modern foundation models often exceed the memory of a single GPU: parameters, activations, optimizer states, and KV caches cannot all reside on one device at useful batch sizes. Even when a model fits, training throughput might be too low without scaling across many GPUs. Parallelism splits the problem along several independent dimensions so that:
Memory is shared across devices (sharding, pipeline stages, sequence splits).
Compute is scaled by processing more data in parallel or by overlapping communication with computation.
Overview of parallelism dimensions#
Dimension |
What is split |
Primary goal |
|---|---|---|
Data parallelism (DP) |
Input batches |
Throughput; same model on each GPU |
FSDP / ZeRO |
Parameters, gradients, optimizer (by stage) |
Memory; keep DP semantics |
Tensor parallelism (TP) |
Individual weight matrices and matmuls |
Memory per layer; needs fast links |
Sequence parallelism (SP) |
Sequence in non-TP regions |
Activation memory with TP |
Pipeline parallelism (PP) |
Layer groups across stages |
Memory; depth-wise split |
Context parallelism (CP) |
Sequence for attention (e.g. ring) |
Very long contexts |
Expert parallelism (EP) |
MoE experts across devices |
Memory and compute for MoE |
These can be combined. The product of parallel degrees must match how processes are laid out on the cluster (see Section 9).
2. Data parallelism (DP)#
In classic data parallelism, every GPU holds a full copy of the model. Each rank receives a different mini-batch of data. After the backward pass, gradients are synchronized so that all ranks apply the same update.
Batch shard 0 Batch shard 1 Batch shard 2 Batch shard 3
| | | |
v v v v
+--------+ +--------+ +--------+ +--------+
| GPU 0 | | GPU 1 | | GPU 2 | | GPU 3 |
| full | | full | | full | | full |
| model | | model | | model | | model |
+--------+ +--------+ +--------+ +--------+
| | | |
+------------------+------------------+------------------+
|
AllReduce(gradients)
|
v
Same weights on all ranks after optimizer step
Properties
Simple to reason about and widely supported.
Requires the full model, activations for one micro-batch, and optimizer state to fit in one GPU’s memory (unless combined with other strategies).
Effective batch size
For a single update that aggregates over data-parallel ranks and gradient accumulation:
[ \text{effective_batch_size} = \text{micro_batch_size} \times \text{num_GPUs}_{\text{DP}} \times \text{gradient_accumulation_steps} ]
Here num_GPUs_DP is the data-parallel group size (not always the same as world_size when TP/PP/EP are also used).
3. Fully sharded data parallel (FSDP / ZeRO)#
ZeRO (Zero Redundancy Optimizer) reduces redundant storage by sharding optimizer states, gradients, and/or parameters across data-parallel ranks.
ZeRO stage |
Sharded |
Idea |
|---|---|---|
Stage 1 |
Optimizer states |
Each rank keeps only (1/N) of optimizer tensors |
Stage 2 |
+ Gradients |
Gradients are sharded; reduced where needed |
Stage 3 |
+ Parameters |
Each rank holds (1/N) of parameters; gather before use |
FSDP (Fully Sharded Data Parallel) in PyTorch is the common implementation of sharded data parallel training; in the Megatron ecosystem, ZeRO-3-style behavior is often discussed alongside FSDP for full parameter sharding.
Typical execution pattern (conceptual)
Forward: AllGather (or equivalent) to materialize parameters needed for the current layer/batch on each rank.
Backward: ReduceScatter (or equivalent) to write shard-sized gradient pieces back to ranks.
Memory intuition
If replicated training used (M) memory per rank for parameters+gradients+optimizer, ideal full sharding across (N) ranks approaches (M/N) for the sharded pieces (plus buffers and fragmentation). Moving from full replication to (1/N) sharding for those tensors saves roughly ((N-1)/N) of that component—for 8 GPUs, about 87.5% of the replicated footprint for the sharded tensors.
In Primus#
Backend |
Configuration |
|---|---|
Megatron-LM |
|
TorchTitan |
|
Exact interactions with checkpoint formats and DDP are documented in Megatron parameters.
4. Tensor parallelism (TP)#
Tensor parallelism splits individual layers (usually linear / attention projections) across GPUs so no single GPU stores the full weight matrix for that layer.
Column-parallel vs row-parallel#
Consider a linear layer (Y = X W) with weight matrix (W). Column-parallel splits (W) along the output dimension (columns). Row-parallel splits (W) along the input dimension (rows) and splits (X) so each rank’s matmul dimensions match.
Column-parallel linear—each rank holds disjoint columns of (W); each rank’s output is a disjoint column shard (half the width for 2-way TP). To recover the full-width tensor the shards are concatenated (All-Gather along the output dim)—this is only done when the full tensor is actually needed:
SAME full X replicated on each TP rank
|
+-----------------+-----------------+
| |
v v
Rank 0: X @ W[:,0:h/2] Rank 1: X @ W[:,h/2:h]
| |
v v
partial Y_0 partial Y_1
(narrow) (narrow)
| |
+-----------------+-----------------+
|
All-Gather (concatenate) on output dim
(only when the full tensor is needed; with
gather_output=False the output stays column-
sharded and feeds the next layer with no comm)
|
v
full-width Y (concatenation of shards)
Row-parallel linear—each rank holds disjoint rows of (W); input (X) is split along the input feature dimension so each rank computes part of the reduction:
Rank 0: X_0 @ W[0:r/2,:] ----+
+-- AllReduce --> Y
Rank 1: X_1 @ W[r/2:r,:] ----+
(X split along features) (partial sums add to full Y)
Typical transformer block pattern: column-parallel for the first projection—its column-sharded output is fed directly into the next layer without communication—then row-parallel for the second projection, which performs the single AllReduce that reconstructs the full output. Column-parallel itself only communicates when gather_output=True.
Communication
Often AllReduce of partial outputs, or ReduceScatter + AllGather sequences depending on implementation and sequence parallelism (see next section).
When to use
Best within a node (NVLink / high-bandwidth GPU–GPU paths). Multi-node TP is possible but latency-sensitive.
In Primus#
Backend |
Parameter |
|---|---|
Megatron-LM |
|
TorchTitan |
|
5. Sequence parallelism (SP)#
Sequence parallelism extends TP by splitting the sequence dimension in regions that are not covered by tensor-parallel matmuls—commonly LayerNorm, dropout, and sometimes residual paths—so activation memory scales better when TP > 1.
Interaction with TP
After a column-parallel region, partial activations can be ReduceScatterd along the sequence.
Before a row-parallel region, activations might be AllGatherd along the sequence.
So SP trades extra collectives for lower per-rank activation footprint on long sequences.
In Primus#
Backend |
Parameter |
|---|---|
Megatron-LM |
|
TorchTitan |
Sequence-parallel behavior is integrated with TP/parallelization pipelines in supported models |
6. Pipeline parallelism (PP)#
Pipeline parallelism assigns disjoint subsets of layers to stages on different devices. Activations (and gradients) move between stages with point-to-point communication.
Microbatch 1: Stage0 -> Stage1 -> Stage2 -> Stage3
Microbatch 2: Stage0 -> Stage1 -> Stage2 -> Stage3
...
Pipeline bubbles#
If a stage waits for input while other stages compute, idle time appears (pipeline bubble). Schedulers reduce bubbles by overlapping forwards and backwards across microbatches.
Common schedules
Schedule |
Idea |
|---|---|
1F1B |
One forward, one backward; classic warmup / steady / cooldown phases |
1F1B interleaved (VPP) |
Virtual pipeline stages: multiple chunks per device to improve utilization |
Zero-bubble (ZB) |
Reorders / splits backward so forward and backward hide each other better; might separate input-gradient vs weight-gradient phases |
V-Schedule / V-Half / V-Min |
Variants reducing bubbles further (names vary by codebase) |
DualPipe |
Bidirectional pipeline scheduling (e.g. DeepSeek-style) to overlap forward/backward paths |
Bubble rate
[ \text{bubble_rate} = \frac{\text{idle time}}{\text{total time}} ]
Lower is better; large microbatch counts and better schedules reduce bubble overhead.
In Primus (Megatron)#
Parameter |
Role |
|---|---|
|
Number of pipeline stages |
|
Enable Primus/Megatron zero-bubble pipeline patches |
|
Use Primus pipeline implementation for schedule logic |
|
e.g. |
See primus/configs/modules/megatron/primus_pipeline.yaml and zero_bubble.yaml in the repo for defaults.
In Primus (TorchTitan)#
Parameter |
Role |
|---|---|
|
Pipeline depth |
|
e.g. |
7. Context parallelism (CP)#
Context parallelism splits the sequence length across devices for long-context training. A common pattern is ring attention: each rank holds a chunk of queries/keys/values and participates in a ring of message passing so attention covers the full sequence without centralizing all activations on one GPU.
Use cases
Long documents, 32K–128K+ tokens, where per-layer activation memory and attention compute must be distributed.
In Primus#
Backend |
Parameter |
|---|---|
Megatron-LM |
|
TorchTitan |
|
8. Expert parallelism (EP)#
Mixture-of-Experts (MoE) models route each token to a small subset of experts. Expert parallelism assigns different experts to different GPUs so expert weights are not duplicated on every device.
Communication
AllToAll (or equivalent) is typical: dispatch tokens to expert ranks and combine expert outputs back.
Expert tensor parallelism (ETP)
Experts can be further tensor-parallel within a subset of GPUs, analogous to TP for dense layers.
In Primus#
Backend |
Parameter |
|---|---|
Megatron-LM |
|
TorchTitan |
|
9. Combining parallelism strategies#
Common pattern#
TP within a node (fast interconnect).
PP across nodes or across groups when layers do not fit on one device.
DP / FSDP for scaling batch size and sharding optimizer state or parameters.
GPU count (simplified)#
For dense models (ignoring CP and detailed MoE layout):
[ \text{world_size} \approx \text{TP} \times \text{PP} \times \text{DP} ]
For MoE-heavy setups, you often see:
[ \text{world_size} \approx \text{TP} \times \text{PP} \times \text{EP} \times \text{DP} ]
Context parallelism introduces another multiplicative factor in layouts where CP ranks are part of the global mesh (exact rank ordering is implementation-specific).
Memory vs communication#
More TP → smaller matrices per GPU but more frequent collectives within layers.
More PP → less memory per stage but pipeline bubbles and latency between stages.
More DP/FSDP → better throughput scaling if communication is not saturated.
Example configurations (illustrative)#
Scenario |
TP |
PP |
DP / notes |
|---|---|---|---|
~7B on 8 GPUs |
1 |
1 |
8-way DP (or FSDP) |
~70B on 64 GPUs (8 nodes × 8) |
8 |
2 |
4-way DP |
Large MoE (e.g. 671B-class) on 256 GPUs |
8 |
4 |
EP 8 (example; real jobs vary widely) |
Always validate against memory profiling, checkpoint sharding, and network on your cluster.
10. Batch size relationships#
Let:
(B_{\text{micro}}) = micro-batch size per forward/backward per data-parallel rank (per step inside accumulation),
(D) = data parallel size (ranks that share the same model split for DP),
(G) = gradient accumulation steps,
(B_{\text{global}}) = global batch size across all DP ranks for one optimizer update.
Then:
[ B_{\text{global}} = B_{\text{micro}} \times D \times G ]
Data parallel size from world size (when using TP, PP, EP):
[ D = \frac{\text{world_size}}{\text{TP} \times \text{PP} \times \text{EP}} ]
(If context parallelism is present, the denominator must include CP in the same way your trainer defines the mesh.)
Solve for accumulation:
[ G = \frac{B_{\text{global}}}{B_{\text{micro}} \times D} ]
Practical notes
Micro batch drives per-GPU activation memory (often linearly in sequence length for attention).
Global batch affects convergence and learning dynamics; scaling laws often refer to global batch.
Gradient accumulation increases time per optimizer step but reduces memory by using smaller (B_{\text{micro}}).
Megatron-specific names for batch arguments appear in Megatron parameters.