~/wiki

ZeRO Optimizer

Confiance : high
zero-optimizerzero-redundancymemory-optimizationdistributed-trainingoptimizer-statesdeepspeedmicrosoft

Zero Redundancy Optimizer (ZeRO) is an advanced memory optimization technique for distributed training that eliminates memory redundancy by partitioning optimizer states, gradients, and parameters across devices while maintaining training efficiency.

Core Problem: Memory Redundancy

Traditional Data Parallelism Issues

In standard distributed training:

  • Each GPU maintains complete copy of model parameters
  • Each GPU stores full optimizer states (often 2-3x parameter size)
  • Each GPU accumulates complete gradient set
  • Result: Massive memory redundancy across devices

Memory Components

For a model with P parameters using Adam optimizer:

  • Model parameters: P values
  • Gradients: P values
  • Optimizer states: 2P values (momentum + variance)
  • Total per GPU: 4P values × number of GPUs

ZeRO Stages

Stage 1: Optimizer State Partitioning

  • Partition: Optimizer states across devices
  • Memory reduction: 4x reduction for Adam optimizer
  • Communication: Gather required states during optimization
  • Benefit: Significant memory savings with minimal overhead

Stage 2: Gradient Partitioning

  • Partition: Gradients in addition to optimizer states
  • Memory reduction: 8x reduction total
  • Communication: All-reduce only assigned gradient partitions
  • Synchronization: Gradients distributed and synchronized efficiently

Stage 3: Parameter Partitioning

  • Partition: Model parameters across devices
  • Memory reduction: Linear with number of devices
  • Communication: Gather parameters as needed for forward/backward
  • Complexity: Most aggressive but requires careful implementation

Implementation Strategy

Dynamic Parameter Management

Stage 3 requires sophisticated parameter handling:

  1. Forward pass: Gather required parameters just before computation
  2. Computation: Execute with temporarily assembled parameters
  3. Cleanup: Discard non-local parameters to free memory
  4. Backward pass: Repeat gathering for gradient computation

Communication Optimization

  • Overlap: Hide parameter gathering with computation
  • Prefetching: Anticipate parameter needs for next layers
  • Bucketing: Group small parameters for efficient communication

Memory Efficiency Gains

Theoretical Reductions

For N devices:

  • Stage 1: Memory per device = (P + P + 2P/N) = (2P + 2P/N)
  • Stage 2: Memory per device = (P + P/N + 2P/N) = (P + 3P/N)
  • Stage 3: Memory per device = (P/N + P/N + 2P/N) = 4P/N

Practical Benefits

  • Larger models: Train models that wouldn't fit in aggregate GPU memory
  • Bigger batches: Use memory savings for increased batch sizes
  • Longer sequences: Handle extended context lengths
  • More devices: Scale to larger numbers of GPUs effectively

Communication Patterns

All-Gather Operations

  • Frequency: Parameter gathering before each layer computation
  • Size: Only required parameter subset
  • Optimization: Overlap with computation when possible

All-Reduce for Gradients

  • Stage 1 & 2: Traditional gradient synchronization
  • Stage 3: Reduced communication volume due to partitioning
  • Bucketing: Efficient handling of small gradient groups

Trade-offs and Considerations

Communication Overhead

  • Increased frequency: More communication operations per training step
  • Network sensitivity: Performance heavily dependent on interconnect bandwidth
  • Latency impact: Higher communication latency affects training speed

Implementation Complexity

  • Stage progression: Each stage adds implementation complexity
  • Memory management: Sophisticated dynamic allocation required
  • Debugging difficulty: Distributed state makes debugging challenging

Framework Integration

DeepSpeed Implementation

  • Native support: ZeRO is core feature of Microsoft's DeepSpeed
  • Automatic optimization: Framework handles communication scheduling
  • Configuration: Simple parameter selection for different stages

Other Framework Support

  • PyTorch FSDP: Similar concepts in Fully Sharded Data Parallel
  • FairScale: Facebook's implementation of sharding strategies
  • Custom implementations: Framework-agnostic manual implementation possible

Performance Optimization

Stage Selection Strategy

Choose optimal stage based on:

  • Memory pressure: How severely memory constrained
  • Network bandwidth: Available inter-device communication
  • Model size: Larger models benefit more from aggressive stages
  • Batch size requirements: Memory needs for target batch size

Hybrid Approaches

  • Selective partitioning: Partition only specific components
  • Gradient accumulation: Combine with micro-batching strategies
  • Mixed precision: Coordinate with FP16/BF16 optimizations

Advanced Optimizations

ZeRO-Offload

  • CPU offloading: Move optimizer states to CPU memory
  • Heterogeneous memory: Utilize both GPU and CPU memory hierarchies
  • Bandwidth management: Balance GPU-CPU transfer costs

ZeRO-Infinity

  • NVMe integration: Use high-speed storage for parameter swapping
  • Memory hierarchy: GPU → CPU → NVMe memory management
  • Extremely large models: Train models larger than total system memory

See also