MuonAdamW vs DistMuonAdamW: Key Differences for Single and Multi-GPU Training in nanochat

The primary difference between MuonAdamW and DistMuonAdamW is that MuonAdamW runs on a single GPU or CPU without inter-process communication, while DistMuonAdamW implements a ZeRO-2-sharded, asynchronous distributed optimizer designed for multi-GPU training without PyTorch DDP overhead.

Both optimizer classes reside in nanochat/optim.py and implement the same hybrid optimization strategy—applying AdamW to regular parameters and Muon (momentum-orthogonalized updates) to 2-D matrix parameters. While they share identical mathematical formulations and constructor signatures, they differ fundamentally in execution context, memory sharding, and communication patterns.

Core Architectural Differences

Execution Context and Communication Strategy

MuonAdamW operates strictly in a single-device context (single GPU or CPU) with no distributed communication. In contrast, DistMuonAdamW (defined at line 297 in nanochat/optim.py) implements a three-phase asynchronous pipeline:

  1. Launch async reduce operations (all-reduce or reduce-scatter)
  2. Wait for reduces, compute updates, and launch gather operations
  3. Wait for gathers and copy results back to parameters

This design overlaps communication with computation, minimizing idle GPU time during distributed training.

Memory Management and State Sharding

MuonAdamW keeps full optimizer state—including exp_avg, exp_avg_sq, momentum_buffer, and second_momentum_buffer—replicated on the single device.

DistMuonAdamW employs ZeRO-2-style sharding: small AdamW parameters remain replicated, while large AdamW parameters have sharded exp_avg and exp_avg_sq tensors. Muon buffers are also sharded per-rank, with each process storing only the slice of momentum_buffer and second_momentum_buffer that it owns. This reduces per-GPU memory usage in proportion to the world size.

Parameter Processing and Padding

While MuonAdamW processes all parameters as-is, DistMuonAdamW must handle distributed tensor slicing. The optimizer stacks Muon groups into a (K, …) tensor, pads the total size to a multiple of the world size, and performs reduce-scatter operations so each rank receives only its contiguous chunk. After the optimization step, padding is stripped following the all-gather phase.

Kernel Invocation and Implementation Details

Both optimizers call the same fused CUDA kernels—adamw_step_fused and muon_step_fused—but DistMuonAdamW wraps them with asynchronous communication primitives including dist.reduce_scatter_tensor, dist.all_reduce, and dist.all_gather_into_tensor. The distributed implementation maintains temporary buffers (stacked_grads, grad_chunk, updated_params) that are reused across communication phases to minimize memory allocation overhead.

Implementation in nanochat/optim.py

MuonAdamW begins at line 152 of nanochat/optim.py and serves as the reference single-GPU implementation. DistMuonAdamW begins at line 297 and extends the base functionality with distributed training machinery.

Both classes accept the same param_groups format with "kind": "adamw" or "kind": "muon" specifications, but DistMuonAdamW automatically handles gradient synchronization and parameter sharding across processes without requiring torch.nn.parallel.DistributedDataParallel.

Usage Examples

Single-GPU Training with MuonAdamW

When training on a single device (as used in scripts/base_train.py when ddp=False), instantiate MuonAdamW directly:

import torch
from nanochat.optim import MuonAdamW

param_groups = [
    {"params": [model.linear.weight], "kind": "muon",
     "lr": 1e-3, "momentum": 0.9, "ns_steps": 5, "beta2": 0.999, "weight_decay": 1e-2},
    {"params": [model.bias], "kind": "adamw",
     "lr": 1e-3, "betas": (0.9, 0.999), "eps": 1e-8, "weight_decay": 1e-2},
]

optimizer = MuonAdamW(param_groups)

for batch in loader:
    loss = model(batch).loss
    loss.backward()
    optimizer.step()
    optimizer.zero_grad()

This implementation runs entirely on the local GPU without torch.distributed calls.

Multi-GPU Distributed Training with DistMuonAdamW

For multi-GPU setups (selected automatically in nanochat/gpt.py when ddp=True), use DistMuonAdamW:

import torch.distributed as dist
from nanochat.optim import DistMuonAdamW

dist.init_process_group(backend="nccl")
torch.cuda.set_device(local_rank)

optimizer = DistMuonAdamW(param_groups)

for batch in loader:
    loss = model(batch).loss
    loss.backward()
    optimizer.step()  # Async reduce → compute → async gather

    optimizer.zero_grad()

Behind the scenes, this launches asynchronous all_reduce operations for AdamW parameters and reduce_scatter/all_gather for Muon parameters, overlapping communication with the compute kernels.

When to Use Each Optimizer

Choose MuonAdamW when training on a single GPU or CPU, debugging distributed logic, or running reference implementations where memory replication is not a constraint. It provides simplicity and minimal overhead for single-device training.

Choose DistMuonAdamW when scaling to multiple GPUs where you need ZeRO-2 memory sharding, asynchronous communication overlap, and the ability to train large models without the full overhead of PyTorch DDP. This optimizer is essential for memory-efficient training across many GPUs in the nanochat framework.

Summary

  • MuonAdamW (line 152 in nanochat/optim.py) provides a single-GPU reference implementation with full state replication and no communication overhead.
  • DistMuonAdamW (line 297 in nanochat/optim.py) adds distributed training capabilities with ZeRO-2 sharding, three-phase async communication, and memory-efficient buffer management.
  • Both optimizers use the same fused kernels (adamw_step_fused, muon_step_fused) and accept identical param_groups structures.
  • DistMuonAdamW handles padding, chunking, and tensor slicing automatically for distributed Muon updates.
  • Selection typically occurs automatically in nanochat/gpt.py based on the ddp flag.

Frequently Asked Questions

Can I use DistMuonAdamW on a single GPU?

While technically possible, DistMuonAdamW requires initializing a distributed process group (dist.init_process_group) even with world size 1, and it introduces unnecessary communication overhead. For single-GPU training, MuonAdamW is the recommended and default choice as implemented in scripts/base_train.py when ddp=False.

How does ZeRO-2 sharding reduce memory in DistMuonAdamW?

According to the implementation in nanochat/optim.py, DistMuonAdamW shards the AdamW optimizer states (exp_avg, exp_avg_sq) across ranks for large parameters, keeping only 1/world_size of the state on each GPU. Muon momentum buffers are similarly partitioned, reducing per-GPU memory usage from O(model_size) to O(model_size / world_size) for optimizer states.

Do both optimizers support identical hyperparameters?

Yes, both classes share the same constructor signature and param_groups format. You can specify lr, betas, eps, and weight_decay for AdamW parameters, and lr, momentum, ns_steps, and beta2 for Muon parameters identically in both optimizers. The distributed version handles the synchronization of these hyperparameter updates transparently across processes.

Why does DistMuonAdamW use async communication instead of standard PyTorch DDP?

The nanochat implementation specifically avoids full DDP overhead by manually launching async dist.reduce_scatter_tensor, dist.all_reduce, and dist.all_gather_into_tensor operations. This three-phase approach (reduce → compute → gather) allows the optimizer to overlap communication with computation and implement ZeRO-2-style sharding that standard DDP does not provide, resulting in better memory efficiency and higher throughput across many GPUs.

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 →