# How the Distributed Optimizer in MegaDLMs Improves Training Performance

> Discover how the distributed optimizer in MegaDLMs boosts training performance by sharding states and overlapping communication for faster, more efficient GPU usage.

- Repository: [Jinjie Ni/megadlms](https://github.com/jinjieni/megadlms)
- Tags: deep-dive
- Published: 2026-03-04

---

**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`](https://github.com/jinjieni/megadlms/blob/main/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`](https://github.com/jinjieni/megadlms/blob/main/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`](https://github.com/jinjieni/megadlms/blob/main/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`](https://github.com/jinjieni/megadlms/blob/main/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`](https://github.com/jinjieni/megadlms/blob/main/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`](https://github.com/jinjieni/megadlms/blob/main/megatron/training/yaml_arguments.py) (lines 139-143):

```python
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:

```python
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`](https://github.com/jinjieni/megadlms/blob/main/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:

```python

# 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`](https://github.com/jinjieni/megadlms/blob/main/megatron/core/optimizer/optimizer_config.py).
- Checkpointing is handled by `sharded_param_state_dp_zero` in [`megatron/core/optimizer/distrib_optimizer.py`](https://github.com/jinjieni/megadlms/blob/main/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`](https://github.com/jinjieni/megadlms/blob/main/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`](https://github.com/jinjieni/megadlms/blob/main/megatron/core/optimizer/distrib_optimizer.py), with configuration flags defined in [`megatron/core/optimizer/optimizer_config.py`](https://github.com/jinjieni/megadlms/blob/main/megatron/core/optimizer/optimizer_config.py). The asynchronous buffer management is handled by `_ParamAndGradBuffer` in [`megatron/core/distributed/param_and_grad_buffer.py`](https://github.com/jinjieni/megadlms/blob/main/megatron/core/distributed/param_and_grad_buffer.py), and the training-loop integration appears in [`megatron/training/training.py`](https://github.com/jinjieni/megadlms/blob/main/megatron/training/training.py) around lines 647-653.