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

> Explore MuonAdamW vs DistMuonAdamW for nanochat training. Understand single GPU vs. multi-GPU optimizations and ZeRO-2 sharding benefits without PyTorch DDP overhead.

- Repository: [Andrej/nanochat](https://github.com/karpathy/nanochat)
- Tags: deep-dive
- Published: 2026-03-10

---

**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`](https://github.com/karpathy/nanochat/blob/main/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`](https://github.com/karpathy/nanochat/blob/main/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`](https://github.com/karpathy/nanochat/blob/main/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`](https://github.com/karpathy/nanochat/blob/main/scripts/base_train.py) when `ddp=False`), instantiate **MuonAdamW** directly:

```python
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`](https://github.com/karpathy/nanochat/blob/main/nanochat/gpt.py) when `ddp=True`), use **DistMuonAdamW**:

```python
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`](https://github.com/karpathy/nanochat/blob/main/nanochat/optim.py)) provides a single-GPU reference implementation with full state replication and no communication overhead.
- **DistMuonAdamW** (line 297 in [`nanochat/optim.py`](https://github.com/karpathy/nanochat/blob/main/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`](https://github.com/karpathy/nanochat/blob/main/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`](https://github.com/karpathy/nanochat/blob/main/scripts/base_train.py) when `ddp=False`.

### How does ZeRO-2 sharding reduce memory in DistMuonAdamW?

According to the implementation in [`nanochat/optim.py`](https://github.com/karpathy/nanochat/blob/main/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.