All articles Distributed Training

Distributed Training at Scale: Megatron-LM 3D Parallelism & DeepSpeed ZeRO-3

Scaling LLM training across 10,000+ GPUs: Tensor Parallelism (TP), Pipeline Parallelism (PP), Sequence Parallelism (SP), and ZeRO-3 Memory Offloading.

Training frontier models with hundreds of billions of parameters (such as Llama 3 405B) exceeds the VRAM capacity of any individual GPU by orders of magnitude. 3D Parallelism — combining Tensor Parallelism (TP), Pipeline Parallelism (PP), and Data Parallelism (DP) with DeepSpeed ZeRO-3 — orchestrates multi-node GPU clusters into a unified compute fabric capable of linear scaling across tens of thousands of accelerators.

1D Data Parallelism replicates models; 3D Parallelism shards model weights, gradients, optimizer states, and activations across entire GPU clusters.

GPU memory consumption breakdown

During FP16 / BF16 mixed-precision training with AdamW, memory consumption per parameter P is allocated as follows:

Model Parameters (16-bit):      2 * P bytes
Gradients (16-bit):             2 * P bytes
AdamW Optimizer States (FP32): 12 * P bytes (4B FP32 copy + 4B momentum + 4B variance)
Total Static Memory:            16 * P bytes

For a 70B parameter model: Static Memory = 16 * 70B = 1,120 GB VRAM!
(Excludes dynamic activation memory)

Megatron Tensor Parallelism (Column & Row GEMM)

Shoeybi et al. (2019) introduced Tensor Parallelism (TP) to split individual matrix multiplications across GPUs within the same NVLink node. In Transformer MLPs:

Column Parallel GEMM (Weight W split column-wise):
Y1 = X * W1,  Y2 = X * W2
Output Y = [Y1, Y2]  (Concatenated without communication)

Row Parallel GEMM (Weight W split row-wise):
Z1 = Y1 * W1_row,  Z2 = Y2 * W2_row
Output Z = AllReduce_Sum(Z1 + Z2)  (1 All-Reduce per Transformer layer)

DeepSpeed ZeRO-1, ZeRO-2, and ZeRO-3 memory sharding

Microsoft's ZeRO (Zero Redundancy Optimizer) eliminates memory redundancies across Data Parallel ranks:

  • ZeRO-1 (Optimizer State Partitioning): Shards AdamW optimizer states across N_dp ranks (4x memory reduction).
  • ZeRO-2 (Gradient Partitioning): Shards gradients alongside optimizer states (8x memory reduction).
  • ZeRO-3 (Parameter Partitioning): Shards model parameters themselves across all Data Parallel ranks. Parameters are gathered dynamically via All-Gather before forward/backward execution and discarded immediately after.

PyTorch Distributed Data Parallel (DDP) & ZeRO-3 config setup

# DeepSpeed ZeRO-3 Configuration JSON
deepspeed_config = {
    "train_batch_size": 128,
    "gradient_accumulation_steps": 4,
    "fp16": {
        "enabled": True
    },
    "zero_optimization": {
        "stage": 3,
        "offload_optimizer": {
            "device": "cpu",
            "pin_memory": True
        },
        "offload_param": {
            "device": "cpu",
            "pin_memory": True
        },
        "overlap_comm": True,
        "allgather_bucket_size": 5e8,
        "reduce_bucket_size": 5e8
    }
}

Cluster communication efficiency

  1. Keep Tensor Parallelism Intra-Node: Enforce TP boundaries within NVLink domain (max 8 GPUs) to avoid slow InfiniBand inter-node latency.
  2. Use Sequence Parallelism (SP): Combine TP with Sequence Parallelism to split Dropout and LayerNorm activations across ranks.
  3. ZeRO-3 Offloading for Massive Models: Offload optimizer states to host CPU RAM when GPU memory boundaries are saturated.
← Back to all articles