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:
- Forward pass: Gather required parameters just before computation
- Computation: Execute with temporarily assembled parameters
- Cleanup: Discard non-local parameters to free memory
- 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
- memory-optimization
- distributed-training
- [[DeepSpeed