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:

  1. Optimizer-state memory explosion: Memory consumption grows as O(#parameters) per GPU, quickly exceeding VRAM limits.
  2. 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:

  1. Async start: Before optimizer.step() executes, the training loop starts an asynchronous all-gather for the first bucket of parameters (see megatron/training/training.py lines 647-653). The handle is stored in _ParamAndGradBuffer.param_gather_handle.
  2. Compute overlap: The optimizer proceeds with the compute-heavy update step while the network transfer proceeds in the background.
  3. Synchronization: After the step completes, finish_param_sync ensures the gather has finished (see megatron/core/distributed/param_and_grad_buffer.py lines 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_optimizer and overlap_param_gather_with_optimizer_step in megatron/core/optimizer/optimizer_config.py.
  • Checkpointing is handled by sharded_param_state_dp_zero in megatron/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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →