How nanochat Combines Muon and AdamW in a Single Optimizer
nanochat merges Muon and AdamW by partitioning parameters into distinct groups—applying AdamW to 0-D and 1-D tensors while using Muon for 2-D matrices—then dispatching to specialized fused kernels within a unified step() method.
The karpathy/nanochat repository implements a hybrid optimization strategy that leverages the complementary strengths of both Muon and AdamW. This combined approach assigns each algorithm to the parameter shapes where it excels most, delivering efficient training for large language models. Understanding how nanochat's optimizer combines Muon and AdamW requires examining its parameter grouping logic, fused kernel implementation, and distributed training support.
Parameter Grouping Strategy
nanochat's optimization begins with intelligent parameter partitioning. The model's setup_optimizer method categorizes every parameter into one of two groups based on dimensionality: tensors with fewer than two dimensions are flagged for AdamW, while 2-D matrices are designated for Muon processing.
This grouping occurs in the training scripts, where the model factory method constructs the optimizer instance. The method accepts separate learning rate arguments for each parameter type, allowing fine-grained control over the optimization dynamics. Referencing the source in [scripts/base_train.py](https://github.com/karpathy/nanochat/blob/master/scripts/base_train.py#L306-L315), the setup call passes distinct learning rates for embeddings, scalars, and matrices, which the optimizer uses to configure its internal parameter groups.
The MuonAdamW Class Architecture
At the core of nanochat's optimization engine sits the MuonAdamW class, a custom PyTorch optimizer that orchestrates the dual-algorithm approach. Unlike traditional optimizers that apply a single update rule universally, this class maintains separate state management and step functions for each parameter category.
The class inherits from torch.optim.Optimizer and overrides the step() method to implement conditional dispatch. When invoked, it iterates through self.param_groups and routes each group to either _step_adamw or _step_muon based on the group's kind attribute. This architectural decision keeps the two optimization paths completely independent while presenting a unified interface to the training loop.
AdamW Step Implementation
The AdamW path handles all non-matrix parameters including embeddings, biases, and layer normalization statistics. The _step_adamw method initializes per-parameter state dictionaries containing first and second moment estimates (exp_avg and exp_avg_sq).
To maximize throughput, the implementation populates several 0-D CPU tensors with current hyperparameters—step count, learning rate, betas, epsilon, and weight decay—then invokes the adamw_step_fused kernel. This fused operation, compiled with @torch.compile, executes the entire AdamW update computation in a single graph, eliminating Python interpreter overhead during the parameter update phase.
Muon Step Implementation
For 2-D weight matrices, nanochat employs the Muon optimizer via the _step_muon method. This implementation handles the orthogonalization-based updates that characterize Muon's approach to optimization. The method aggregates gradients and parameters from all matrix-shaped tensors in the group into stacked tensors to enable vectorized processing.
The Muon step maintains specialized momentum buffers and performs Polar-Express orthogonalization along with variance reduction techniques. After preparing the 0-D hyperparameter tensors for Muon-specific settings, it calls the muon_step_fused kernel. This compiled kernel applies Nesterov momentum, performs the orthogonalization, and applies cautious weight decay before writing the updated values back to the original parameter tensors.
Fused Kernel Optimization
Both optimization paths rely on custom fused kernels to minimize dispatch overhead. The adamw_step_fused and muon_step_fused functions are decorated with @torch.compile, allowing PyTorch to generate optimized machine code for the entire update computation.
This design choice proves critical for training efficiency, as it prevents the Python global interpreter lock from becoming a bottleneck during parameter updates. By compiling the inner loops into optimized kernels, nanochat achieves performance comparable to manually written CUDA operations while maintaining the flexibility of Pythonic optimizer implementations.
Distributed Training Support
When training across multiple GPUs with Distributed Data Parallel (DDP), nanochat substitutes MuonAdamW with DistMuonAdamW. This subclass implements a three-phase asynchronous communication pattern that overlaps gradient reduction with computation.
The distributed variant maintains the same dual-algorithm structure but adds specialized communication hooks. These hooks coordinate all-reduce operations for both AdamW and Muon parameter groups while preserving the independent update semantics of the single-GPU implementation. The result allows linear scaling of the hybrid optimization strategy across many accelerators without modifying the underlying training loop code.
Practical Usage Examples
Instantiating the combined optimizer requires a single call to the model's setup method. The following pattern appears throughout nanochat's training scripts:
# Initialize the hybrid optimizer with separate learning rates
optimizer = model.setup_optimizer(
unembedding_lr=args.unembedding_lr * batch_lr_scale,
embedding_lr=args.embedding_lr * batch_lr_scale,
scalar_lr=args.scalar_lr * batch_lr_scale,
matrix_lr=args.matrix_lr * batch_lr_scale,
weight_decay=weight_decay_scaled,
)
Within the training loop, usage remains identical to standard PyTorch optimizers:
for iteration, (inputs, targets) in enumerate(train_loader):
optimizer.zero_grad(set_to_none=True)
outputs = model(inputs)
loss = compute_loss(outputs, targets)
loss.backward()
optimizer.step() # Dispatches to AdamW or Muon based on parameter shape
For distributed configurations, the model factory automatically returns a DistMuonAdamW instance when args.ddp is enabled, requiring no changes to the optimization loop:
if args.ddp:
optimizer = model.setup_optimizer(...) # Returns DistMuonAdamW
else:
optimizer = model.setup_optimizer(...) # Returns MuonAdamW
Summary
- nanochat partitions parameters into two distinct groups: 0-D/1-D tensors use AdamW while 2-D matrices use Muon.
- The
MuonAdamWclass dispatches to specialized step methods based on thekindattribute of each parameter group. - Fused kernels
adamw_step_fusedandmuon_step_fusedcompile the update logic for maximum performance. - The
DistMuonAdamWvariant extends this architecture to multi-GPU training with asynchronous communication.
Frequently Asked Questions
Which parameters does Muon optimize versus AdamW in nanochat?
Muon exclusively handles 2-D matrix-shaped weight parameters, such as the weight matrices in linear and attention layers. AdamW manages all other parameters including 1-D embeddings, biases, layer normalization statistics, and scalar values, ensuring stable optimization for non-matrix tensors.
Why does nanochat use fused kernels for the optimizer steps?
The fused kernels eliminate Python interpreter overhead by compiling the entire parameter update computation into optimized machine code using @torch.compile. This approach maximizes GPU utilization and training throughput, particularly important when updating billions of parameters in large language models.
How does distributed training work with the combined optimizer?
The DistMuonAdamW class extends the base optimizer with a three-phase asynchronous communication pattern that handles gradient all-reduction across GPUs. It preserves the dual-algorithm structure while overlapping communication with computation, enabling efficient scaling without code changes to the training loop.
Can I adjust learning rates independently for Muon and AdamW parameters?
Yes, the setup_optimizer method accepts separate learning rate arguments including matrix_lr for Muon-optimized weights and distinct rates for embeddings, scalars, and unembedding layers. This design allows precise control over the optimization dynamics for different parameter types within the same training run.
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 →