How the Distributed Optimizer in MegaDLMs Improves Training Performance
The distributed optimizer in MegaDLMs improves training by sharding optimizer states across data-parallel ranks to reduce GPU memory usage, while overlapping parameter all-gather operations with the optimizer computation step to eliminate synchronous communication bottlenecks.
MegaDLMs (Megatron-based Deep Language Models) scale to billions of parameters across hundreds of GPUs, where traditional data-parallel training hits severe memory and communication walls. The distributed optimizer, implemented in the jinjieni/megadlms repository, replaces the naive replicate-everywhere approach with a sharded, asynchronous design that enables larger models and higher throughput.
What Problem Does the Distributed Optimizer Solve?
Standard data-parallel (DP) training stores a full copy of every optimizer tensor—such as Adam’s exp_avg and exp_avg_sq—on each DP rank. With multi-billion-parameter models, this creates two critical bottlenecks:
- Optimizer-state memory explosion: Memory consumption grows as O(#parameters) per GPU, quickly exceeding VRAM limits.
- Synchronous communication overhead: After each backward pass, ranks must all-reduce full gradients and broadcast updated parameters, stalling the compute pipeline while waiting for network transfers.
The distributed optimizer solves both issues through state sharding and communication-computation overlap.
How Sharding Optimizer State Reduces Memory
When use_distributed_optimizer=True is set in OptimizerConfig (see megatron/core/optimizer/optimizer_config.py lines 108-112), the optimizer constructs a _ParamAndGradBuffer that splits contiguous gradients and parameters into DP-world-size shards. Each rank owns only the shards intersecting its data partition.
The mapping from model parameters to their respective shards is built by DistributedOptimizer._build_model_gbuf_param_range_map in megatron/core/optimizer/distrib_optimizer.py lines 95-126. This architecture fundamentally changes resource requirements:
| Aspect | Naïve DP | Distributed Optimizer |
|---|---|---|
| Optimizer-state memory per GPU | Full copy → O(#params) | Sharded → O(#params / DP-world-size) |
| State checkpointing | All ranks write identical tensors | Only rank 0 writes sharded state (see sharded_param_state_dp_zero at lines 1281-1304) |
| Parameter-update traffic | All-reduce over full tensors | Reduce-scatter on shards → less traffic |
By sharding optimizer tensors, models that previously required >80 GB of GPU memory can now train on a single 40 GB GPU when using 8 DP ranks.
Overlapping Communication with Computation
The distributed optimizer eliminates the synchronous parameter-gather stall through asynchronous communication overlap.
Asynchronous Parameter Gathering
MegaDLMs expose the flag overlap_param_gather_with_optimizer_step in OptimizerConfig (see megatron/core/optimizer/optimizer_config.py lines 114-116). When enabled:
- Async start: Before
optimizer.step()executes, the training loop starts an asynchronous all-gather for the first bucket of parameters (seemegatron/training/training.pylines 647-653). The handle is stored in_ParamAndGradBuffer.param_gather_handle. - Compute overlap: The optimizer proceeds with the compute-heavy update step while the network transfer proceeds in the background.
- Synchronization: After the step completes,
finish_param_syncensures the gather has finished (seemegatron/core/distributed/param_and_grad_buffer.pylines 70-84).
This overlap converts a serial compute-communication sequence into a pipeline, reducing per-step wall-clock time by approximately 10–15% on typical 8-node configurations.
Configuration and Implementation
Enabling the Distributed Optimizer
Configure the optimizer through YAML or command-line arguments passed via megatron/training/yaml_arguments.py (lines 139-143):
from megatron.training.yaml_arguments import parse_yaml_args
args = parse_yaml_args('configs/gpt2.yaml')
# Enable distributed optimizer and communication overlap
args.use_distributed_optimizer = True
args.overlap_param_gather = True
args.overlap_param_gather_with_optimizer_step = True
# The optimizer is instantiated via get_megatron_optimizer() which
# returns DistributedOptimizer when the flag is set
Inspecting Sharded Optimizer States
During training, only rank 0 maintains the sharded checkpoint:
if torch.distributed.get_rank() == 0:
# Access sharded state via state_dict()
sharded_state = optimizer.state_dict()
torch.save(sharded_state, 'ckpt/optimizer_shard_rank0.pt')
This pattern follows the implementation in megatron/core/optimizer/distrib_optimizer.py lines 1281-1304, where sharded_param_state_dp_zero handles rank-exclusive serialization.
Verifying Overlap Behavior
Unit tests in the repository confirm the asynchronous handle creation:
# tests/unit_tests/distributed/test_param_and_grad_buffer.py
def test_overlap_param_gather():
args.overlap_param_gather = True
args.use_distributed_optimizer = True
buffer.start_param_sync()
# Assert async handle was created for overlapped gathering
assert buffer.param_gather_handle is not None
Performance Impact and Scalability
The distributed optimizer delivers measurable gains across three dimensions:
- Memory efficiency: Linear reduction in optimizer-state memory per GPU proportional to DP world size, enabling larger models on commodity hardware.
- Throughput: Overlapping the first-bucket gather saves 10–15% of step time by hiding communication latency behind computation.
- Scalability: Because each rank only reduces and scatters its own shard, collective communication volume scales linearly with DP ranks rather than remaining constant, maintaining efficiency up to dozens of nodes.
Summary
- The distributed optimizer shards Adam states and momentum buffers across data-parallel ranks, reducing per-GPU memory from O(#params) to O(#params / DP-world-size).
- Communication overlap is achieved by starting the first-bucket parameter all-gather asynchronously before the optimizer step, hiding network latency.
- Key configuration occurs through
OptimizerConfig.use_distributed_optimizerandoverlap_param_gather_with_optimizer_stepinmegatron/core/optimizer/optimizer_config.py. - Checkpointing is handled by
sharded_param_state_dp_zeroinmegatron/core/optimizer/distrib_optimizer.py, ensuring only rank 0 writes to disk.
Frequently Asked Questions
What is the main advantage of the distributed optimizer in MegaDLMs?
The primary advantage is the reduction of GPU memory usage through sharding, allowing models that previously required >80 GB of VRAM to train on 40 GB GPUs when using 8 data-parallel ranks, while also improving speed by overlapping communication with computation.
How does the distributed optimizer handle checkpointing?
Only rank 0 writes the optimizer state to disk via the sharded_param_state_dp_zero method in distrib_optimizer.py (lines 1281-1304). The state_dict() method returns the sharded view, so each rank stores only its portion of the Adam momentum and variance tensors.
Can I use the distributed optimizer without overlapping communication?
Yes. You can set use_distributed_optimizer=True while keeping overlap_param_gather_with_optimizer_step=False. This still provides memory savings from sharding, but sacrifices the 10–15% step-time improvement that comes from hiding the parameter all-gather latency behind the optimizer computation.
Where is the distributed optimizer implemented in the codebase?
The core logic resides in megatron/core/optimizer/distrib_optimizer.py, with configuration flags defined in megatron/core/optimizer/optimizer_config.py. The asynchronous buffer management is handled by _ParamAndGradBuffer in megatron/core/distributed/param_and_grad_buffer.py, and the training-loop integration appears in megatron/training/training.py around lines 647-653.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →