TempatGunting

Distributed Training with DeepSpeed and FSDP: Scaling LLMs

Training a 7B parameter model on a single GPU is a non-starter — the model weights alone consume 14GB in float16, and optimizer states triple that to 42GB before you even account for activations and gradient buffers. Distributed training frameworks solve this by splitting model state across multiple devices, but the two dominant options — Microsoft's DeepSpeed and PyTorch's FSDP — make different tradeoffs that matter in production. After running both across clusters ranging from 8 to 256 GPUs for pre-training and fine-tuning workloads, here's what you need to know to pick the right one.

The Memory Problem

A mixed-precision training run for a model with Ψ parameters using the AdamW optimizer stores four copies of each parameter: fp16 model weights (2Ψ bytes), fp16 gradients (2Ψ bytes), fp32 master weights (4Ψ bytes), fp32 momentum (4Ψ bytes), and fp32 variance (4Ψ bytes). That's 16Ψ bytes of model state per parameter — 112GB for a 7B model, before activations.

Standard data parallelism replicates all of this on every GPU. That's wasteful: each GPU holds a full copy of optimizer states it will never update independently. The insight behind both ZeRO (Zero Redundancy Optimizer) and FSDP is to shard this state across GPUs, so each device stores only its fraction. The frameworks differ in how they implement the sharding, what additional optimizations they offer, and how they integrate with the rest of your training stack. Understanding transformer architecture helps contextualize why these memory demands grow so quickly with model size.

Per-GPU Memory Usage — 7B Model on 8 GPUs DDP ZeRO-1 ZeRO-2 ZeRO-3/FSDP 112 GB 56 GB 32 GB 14 GB 0 GB 28 GB 56 GB 84 GB 112 GB

DeepSpeed ZeRO: Three Stages of Sharding

DeepSpeed's ZeRO optimizer progressively shards more state across GPUs. Each stage builds on the previous one:

ZeRO Stage 1: Optimizer State Partitioning

Each GPU holds only 1/N of the optimizer states (momentum and variance for AdamW). Model parameters and gradients are still fully replicated. This cuts optimizer memory by N× while adding zero communication overhead beyond standard data parallelism. For a 7B model on 8 GPUs, per-device optimizer memory drops from 56GB to 7GB.

ZeRO Stage 2: Gradient Partitioning

In addition to optimizer state sharding, each GPU retains only the gradients corresponding to its optimizer partition. After the backward pass, gradients are reduce-scattered (not all-reduced), so each GPU receives only its slice. This saves 2Ψ/N bytes per device and actually reduces communication volume compared to DDP's all-reduce.

ZeRO Stage 3: Parameter Partitioning

The most aggressive stage: model parameters themselves are sharded. Each GPU stores only 1/N of the parameters and gathers them on-demand during forward and backward passes. This adds two all-gather operations per layer (forward and backward) but enables training models that don't fit in aggregate GPU memory when combined with offloading.

# DeepSpeed ZeRO-3 configuration
ds_config = {
    "train_batch_size": 64,
    "gradient_accumulation_steps": 8,
    "fp16": {"enabled": True},
    "zero_optimization": {
        "stage": 3,
        "offload_optimizer": {
            "device": "cpu",
            "pin_memory": True
        },
        "offload_param": {
            "device": "cpu",
            "pin_memory": True
        },
        "overlap_comm": True,
        "contiguous_gradients": True,
        "sub_group_size": 1e9,
        "reduce_bucket_size": "auto",
        "stage3_prefetch_bucket_size": "auto",
        "stage3_param_persistence_threshold": "auto",
        "stage3_max_live_parameters": 1e9,
        "stage3_max_reuse_distance": 1e9
    }
}

The killer feature of DeepSpeed is ZeRO-Offload and ZeRO-Infinity: offloading optimizer states and parameters to CPU RAM or NVMe storage. This lets you train a 13B model on a single GPU by using system memory as overflow — slow, but functional for fine-tuning workloads where throughput isn't critical. For optimizing batch sizes with gradient accumulation strategies, see our training techniques comparison.

PyTorch FSDP: Native Distributed Training

Fully Sharded Data Parallel (FSDP) is PyTorch's answer to ZeRO-3. Introduced in PyTorch 1.11 and significantly rewritten in PyTorch 2.x, FSDP shards parameters, gradients, and optimizer states across all ranks. The key difference from DeepSpeed: FSDP is a first-class PyTorch citizen, which means it composes naturally with torch.compile, DTensor, and the broader PyTorch ecosystem.

import torch
from torch.distributed.fsdp import (
    FullyShardedDataParallel as FSDP,
    ShardingStrategy,
    MixedPrecision,
    BackwardPrefetch,
)
from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy
from transformers import LlamaForCausalLM, LlamaConfig

model = LlamaForCausalLM(LlamaConfig(
    hidden_size=4096,
    num_hidden_layers=32,
    num_attention_heads=32,
))

mp_policy = MixedPrecision(
    param_dtype=torch.bfloat16,
    reduce_dtype=torch.bfloat16,
    buffer_dtype=torch.bfloat16,
)

auto_wrap_policy = transformer_auto_wrap_policy(
    transformer_layer_cls={LlamaDecoderLayer},
)

model = FSDP(
    model,
    sharding_strategy=ShardingStrategy.FULL_SHARD,
    mixed_precision=mp_policy,
    auto_wrap_policy=auto_wrap_policy,
    backward_prefetch=BackwardPrefetch.BACKWARD_PRE,
    device_id=torch.cuda.current_device(),
    limit_all_gathers=True,
)

FSDP's ShardingStrategy provides equivalent options to ZeRO stages: FULL_SHARD (ZeRO-3), SHARD_GRAD_OP (ZeRO-2), and NO_SHARD (standard DDP). The HYBRID_SHARD strategy shards within a node and replicates across nodes — a practical middle ground that avoids cross-node parameter gathering while still reducing per-GPU memory.

Head-to-Head: Performance Comparison

We benchmarked both frameworks training a 7B Llama-style model on a cluster of 4 nodes, each with 8 A100-80GB GPUs (32 GPUs total), measuring throughput (tokens per second), peak memory, and time to convergence.

MetricDeepSpeed ZeRO-3FSDP Full ShardFSDP Hybrid Shard
Throughput (tokens/s/GPU)3,4203,3803,650
Peak Memory (GB)62.164.371.8
Cross-node Comm (GB/step)28.428.412.6
Step Time (ms)2,3402,3702,190
torch.compile CompatibleLimitedFullFull
Activation CheckpointingBuilt-inBuilt-inBuilt-in
CPU OffloadingFull (optimizer + params)Partial (optimizer only)Partial

Key observations from our benchmarks:

  • Throughput: Nearly identical for full sharding. FSDP Hybrid Shard wins when inter-node bandwidth is the bottleneck (common with 100Gbps InfiniBand), because it eliminates cross-node parameter all-gathers.
  • Memory: DeepSpeed uses slightly less memory due to more aggressive prefetching controls and buffer management. The 2GB difference comes from DeepSpeed's reduce_bucket_size auto-tuning.
  • Compilation: FSDP + torch.compile delivers 10-15% throughput improvement on PyTorch 2.x. DeepSpeed's custom CUDA kernels conflict with torch.compile in some configurations.
  • Offloading: DeepSpeed's CPU/NVMe offloading is production-grade. FSDP supports optimizer state offloading but not parameter offloading with the same maturity.

Communication Patterns and Network Design

Understanding the communication patterns helps you design your cluster networking. Both ZeRO-3 and FSDP Full Shard use three collective operations per training step:

  1. Forward pass all-gather: Each layer gathers its full parameters from all ranks before computing. With prefetching, the next layer's gather overlaps with the current layer's compute.
  2. Backward pass all-gather: Same pattern during backward — parameters are re-gathered for gradient computation.
  3. Backward reduce-scatter: Gradients are reduce-scattered so each rank accumulates only its shard's gradients.

The total communication volume per step is 3Ψ bytes (in the parameter dtype) — three full model transfers. For a 7B model in bf16, that's 42GB per step across all ranks. On an 8-GPU node with NVLink (600 GB/s aggregate), intra-node communication takes ~70ms. Across nodes with 400Gbps InfiniBand, the same transfer takes ~840ms — 12× slower.

This is why FSDP's Hybrid Shard strategy matters: by sharding only within a node and replicating across nodes, it replaces cross-node all-gathers with a single cross-node all-reduce of gradients (2Ψ bytes), cutting cross-node traffic by ~55%. For applications where training data flows through complex orchestration pipelines, optimizing the training communication pattern is equally critical for end-to-end efficiency.

Activation Checkpointing

Even with full parameter sharding, activations dominate GPU memory for long sequences. A 7B model with sequence length 4096 and batch size 1 produces ~16GB of activations. Activation checkpointing (gradient checkpointing) trades compute for memory by recomputing activations during the backward pass instead of storing them.

# FSDP with activation checkpointing
from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import (
    apply_activation_checkpointing,
    checkpoint_wrapper,
    CheckpointImpl,
)

apply_activation_checkpointing(
    model,
    checkpoint_wrapper_fn=checkpoint_wrapper,
    check_fn=lambda submodule: isinstance(submodule, LlamaDecoderLayer),
)

# DeepSpeed activation checkpointing
import deepspeed
deepspeed.checkpointing.configure(
    mpu_=None,
    partition_activations=True,  # Shard activations across GPUs
    contiguous_checkpointing=True,
    checkpoint_in_cpu=False,
)

DeepSpeed offers a unique optimization: partitioned activation checkpointing, which shards activation memory across GPUs rather than replicating it. This reduces per-GPU activation memory by N× at the cost of additional all-gather communication during the backward recomputation. In our testing, partitioned checkpointing adds ~8% training overhead but enables 2× longer sequences. For managing these training experiments at scale, an experiment tracking system becomes essential to compare configurations.

Mixed Precision Strategies

Both frameworks support mixed precision training, but the details matter for numerical stability:

bf16 vs fp16

bf16 (brain floating point) maintains the same dynamic range as fp32 (8 exponent bits) with reduced precision (7 mantissa bits vs 23). fp16 has higher precision (10 mantissa bits) but lower dynamic range (5 exponent bits), requiring loss scaling to prevent gradient underflow. For LLM training, bf16 is almost universally preferred — it eliminates the need for loss scaling and gradient clipping while providing sufficient precision for convergence. Understanding the numerical foundations matters for quantization in production deployment as well.

# FSDP mixed precision — bf16 for everything
mp_policy = MixedPrecision(
    param_dtype=torch.bfloat16,
    reduce_dtype=torch.bfloat16,
    buffer_dtype=torch.bfloat16,
)

# DeepSpeed bf16 config
ds_config = {
    "bf16": {
        "enabled": True
    },
    "zero_optimization": {
        "stage": 3
    }
}

Multi-Node Training: Practical Setup

Distributed training across multiple nodes introduces failure modes that don't exist on a single machine. Here's the production setup we use:

# Launch script for multi-node FSDP training
# Run this on each node, adjusting NODE_RANK
torchrun \
    --nnodes=4 \
    --nproc_per_node=8 \
    --node_rank=$NODE_RANK \
    --master_addr=$MASTER_ADDR \
    --master_port=29500 \
    --rdzv_backend=c10d \
    --rdzv_endpoint=$MASTER_ADDR:29400 \
    train.py \
    --model_name llama-7b \
    --batch_size 2 \
    --gradient_accumulation_steps 4 \
    --learning_rate 3e-4 \
    --warmup_steps 2000 \
    --max_steps 100000 \
    --fsdp_sharding full_shard \
    --activation_checkpointing true

Critical production considerations:

  • NCCL tuning: Set NCCL_IB_DISABLE=0, NCCL_NET_GDR_LEVEL=5 for GPU Direct RDMA, and NCCL_SOCKET_IFNAME to your InfiniBand interface. Wrong NCCL settings can reduce throughput by 50%.
  • Elastic training: DeepSpeed has built-in elastic training support. FSDP relies on torchrun's elastic launch. Both handle node failures by checkpointing and resuming, but DeepSpeed's implementation is more battle-tested.
  • Checkpoint storage: Use async checkpointing to avoid blocking training. DeepSpeed saves sharded checkpoints natively. FSDP checkpoints require StateDictType.FULL_STATE_DICT for portable saves (slow) or SHARDED_STATE_DICT for fast saves (tied to world size).
  • Monitoring: Track GPU utilization, memory, and communication bandwidth per node. A single slow node in a synchronous training setup throttles all nodes. For building monitoring dashboards for your training runs, see our model monitoring guide.

When to Use What

After running both frameworks across dozens of training jobs, here's our decision framework:

Choose DeepSpeed when:

  • You need CPU or NVMe offloading (ZeRO-Offload/Infinity) for models exceeding aggregate GPU memory
  • You're using pipeline parallelism (DeepSpeed's PipelineEngine)
  • You need elastic training with automatic recovery from node failures
  • Your training setup includes heterogeneous hardware (mixed GPU generations)
  • You're using the Hugging Face Trainer — DeepSpeed integration is mature

Choose FSDP when:

  • You want torch.compile optimizations (10-15% throughput gain)
  • You're composing with tensor parallelism via DTensor (2D parallelism)
  • You prefer zero external dependencies beyond PyTorch
  • You're building a custom training loop and want full control over the distributed semantics
  • Your cluster has fast inter-node networking and you can use Hybrid Shard

For fine-tuning workloads using LoRA or QLoRA, both frameworks work well, but DeepSpeed ZeRO-3 with offloading enables fine-tuning 70B+ models on surprisingly modest hardware — 4× A100-40GB with CPU offload can handle QLoRA on a 70B model at ~300 tokens/second.

The landscape is consolidating. PyTorch's FSDP2 (currently in beta) addresses many gaps — including better offloading and composability with torch.compile. Long-term, FSDP's integration with the PyTorch core means it will likely become the default for most teams. But today, DeepSpeed's offloading capabilities and pipeline parallelism still justify its complexity for the largest training runs. For managing the GPU clusters that underpin these training workloads, proper resource management is equally important.

FAQ

What is the difference between DeepSpeed ZeRO and PyTorch FSDP?

Both shard optimizer states, gradients, and parameters across GPUs to reduce per-device memory. DeepSpeed ZeRO offers three explicit stages with advanced CPU/NVMe offloading. FSDP is PyTorch-native, requires no external library, and integrates directly with torch.compile. ZeRO-3 and FSDP full-shard are functionally equivalent in sharding behavior.

When should I use DeepSpeed instead of FSDP?

Use DeepSpeed when you need CPU or NVMe offloading for models exceeding aggregate GPU memory, when training with heterogeneous hardware, or when you want pipeline parallelism via DeepSpeed's PipelineEngine. FSDP is preferred when you want to stay within the PyTorch ecosystem and use torch.compile optimizations.

How much GPU memory does ZeRO-3 save compared to standard data parallelism?

ZeRO-3 shards optimizer states, gradients, and model parameters across all GPUs. For a 7B parameter model with AdamW on 8 GPUs, per-GPU memory drops from ~112GB to ~14GB for model state alone — an 8× reduction.

Can I combine FSDP with tensor parallelism?

Yes. PyTorch 2.x supports composing FSDP with tensor parallelism via DTensor. The recommended approach is TP within a node (across NVLink-connected GPUs) and FSDP across nodes. This 2D parallelism reduces inter-node communication compared to FSDP alone.