# How to Configure AdamConfig and Optimizer Settings in Levanter

> Learn to configure AdamConfig and optimizer settings in Levanter for learning rate schedules, weight decay, gradient clipping, and step skipping. Build Optax optimizers easily.

- Repository: [The Marin Project/marin](https://github.com/marin-community/marin)
- Tags: how-to-guide
- Published: 2026-08-29

---

**Levanter's `AdamConfig` class in [`levanter/optim/config.py`](https://github.com/marin-community/marin/blob/main/levanter/optim/config.py) provides a dataclass-based interface to configure the Adam optimizer with support for learning rate schedules, weight decay masking, gradient clipping, and step skipping, which you instantiate and call `.build(num_train_steps)` to create an Optax optimizer.**

The marin-community/marin repository's Levanter framework implements a type-safe, composable approach to optimizer configuration using Python dataclasses. When you configure AdamConfig and optimizer settings in Levanter, you work with a declarative system that compiles into an Optax gradient transformation pipeline. This architecture separates hyperparameter specification from optimizer construction while providing hooks for advanced regularization techniques like weight decay masking and step skipping.

## Understanding the AdamConfig Architecture

The `AdamConfig` class is defined in [`lib/levanter/src/levanter/optim/config.py`](https://github.com/marin-community/marin/blob/main/lib/levanter/src/levanter/optim/config.py) (lines 57-71) and inherits from the generic `OptimizerConfig` base class. This dataclass declares hyperparameters including `beta1`, `beta2`, `epsilon`, optional gradient-norm clipping (`max_grad_norm`), Nesterov momentum, RMS-clipping, update-norm clipping, step-skipping, and the AdamC corrected weight-decay flag.

```python
from levanter.optim.config import AdamConfig

adam_cfg = AdamConfig(learning_rate=6e-4)          # uses cosine schedule by default

optimizer = adam_cfg.build(num_train_steps=10_000)  # returns an optax.GradientTransformation

```

The configuration object acts as a factory: you first instantiate it with desired hyperparameters, then call `build(num_train_steps)` to compile the configuration into a functional Optax optimizer.

## Configuring Learning Rate Schedules

`OptimizerConfig` includes a `lr_schedule` field (lines 154-179) that accepts either a concrete schedule object or a string name. The schedule is constructed via the `lr_scheduler(num_train_steps)` helper method. Supported schedules include `ConstantLrSchedule`, `CosineLrSchedule`, and other standard decay patterns.

```python
from levanter.optim.config import AdamConfig, CosineLrSchedule

adam_cfg = AdamConfig(
    learning_rate=3e-4,
    lr_schedule=CosineLrSchedule(),               # explicit schedule object

    update_rms_clipping=1.0,                      # RMS clipping on the update

    max_grad_norm=1.0,                            # global-norm clipping

)
optimizer = adam_cfg.build(num_train_steps=5_000)

```

When you pass an explicit schedule object, Levanter bypasses the string-to-object resolution logic and uses your configuration directly.

## Weight Decay and Masking Strategies

`OptimizerConfig` provides two mechanisms for controlling which parameters receive weight decay. First, `weight_decay_modules` accepts a regex pattern or list of module name patterns that identify which parameters should have weight decay applied. Second, `default_weight_decay_mask` is a boolean flag that toggles a sensible default mask when no explicit pattern is provided.

Both mechanisms are applied inside `AdamConfig.build` via the `build_weight_decay_mask` function (lines 95-101) when constructing the weight-decay term in the Optax chain.

```python
adam_cfg = AdamConfig(
    learning_rate=2e-4,
    weight_decay=0.01,
    weight_decay_modules=r".*attention.*weight|.*mlp.*weight",  # regex mask

    adamc_weight_decay=True,    # keep weight_decay / lr constant (AdamC correction)

)
optimizer = adam_cfg.build(num_train_steps=20_000)

```

Setting `adamc_weight_decay=True` enables the corrected weight decay formulation that maintains constant weight decay strength independent of the learning rate.

## Gradient Clipping and Regularization Options

The implementation supports multiple regularization strategies that wrap the core Adam update rule. You can apply global gradient norm clipping via `max_grad_norm`, RMS clipping on updates via `update_rms_clipping`, and configure step skipping for training stability.

Step skipping requires importing `SkipStepConfig` from [`lib/levanter/optim/skipstep.py`](https://github.com/marin-community/marin/blob/main/lib/levanter/optim/skipstep.py), which implements logic to detect and skip optimizer steps when gradients explode beyond a threshold.

```python
from levanter.optim.config import AdamConfig
from levanter.optim.skipstep import SkipStepConfig

adam_cfg = AdamConfig(
    learning_rate=5e-4,
    skip_bad_steps=SkipStepConfig(history_len=128, sigma_factor=6.0)  # custom skip config

)
optimizer = adam_cfg.build(num_train_steps=15_000)

```

## Building the Optimizer with build()

The `AdamConfig.build(num_train_steps)` method (lines 103-150) constructs an Optax transformer pipeline through six distinct stages:

1. **Optional global-norm clipping** via `optax.clip_by_global_norm` when `max_grad_norm` is set
2. **Adam core update rule** via `optax.scale_by_adam` using `beta1`, `beta2`, and `epsilon`
3. **Optional weight decay** via `optax.add_decayed_weights` using the mask built from `weight_decay_modules`
4. **Optional RMS-clipping** or custom update-norm clipping
5. **Learning rate scaling** using the injected schedule
6. **Optional step-skipping wrapper** when `skip_bad_steps` is configured

The method returns a fully configured `optax.GradientTransformation` ready for JIT compilation with JAX.

## Integration in Training Scripts

Training entry points such as [`train_lm.py`](https://github.com/marin-community/marin/blob/main/train_lm.py) and [`train_dpo.py`](https://github.com/marin-community/marin/blob/main/train_dpo.py) expose an `optimizer: OptimizerConfig` field with `default_factory=AdamConfig`, allowing users to override any Adam settings from command line arguments or configuration files (see [`lib/levanter/src/levanter/main/train_lm.py`](https://github.com/marin-community/marin/blob/main/lib/levanter/src/levanter/main/train_lm.py), line 37).

```python
from dataclasses import dataclass, field
from levanter.optim.config import AdamConfig, OptimizerConfig

@dataclass
class TrainConfig:
    optimizer: OptimizerConfig = field(default_factory=AdamConfig)  # default Adam

    # … other training fields …

cfg = TrainConfig()
opt = cfg.optimizer.build(num_train_steps=30_000)

```

This pattern enables type-safe configuration inheritance while allowing runtime overrides through Levanter's configuration resolution system.

## Summary

- **`AdamConfig`** lives in [`lib/levanter/src/levanter/optim/config.py`](https://github.com/marin-community/marin/blob/main/lib/levanter/src/levanter/optim/config.py) (lines 57-71) and inherits from `OptimizerConfig`
- Configure learning rate schedules through the `lr_schedule` field accepting objects like `CosineLrSchedule` or string identifiers
- Control weight decay application using `weight_decay_modules` regex patterns or `default_weight_decay_mask`, with AdamC correction available via `adamc_weight_decay`
- Apply gradient clipping through `max_grad_norm` and `update_rms_clipping` parameters
- Build functional optimizers by calling `.build(num_train_steps)`, which returns an `optax.GradientTransformation`
- Integrate into training scripts using dataclass fields with `default_factory=AdamConfig` to enable CLI overrides

## Frequently Asked Questions

### How do I enable gradient clipping in Levanter's AdamConfig?

Set the `max_grad_norm` parameter to a positive float value when instantiating `AdamConfig`. According to the source code in [`lib/levanter/src/levanter/optim/config.py`](https://github.com/marin-community/marin/blob/main/lib/levanter/src/levanter/optim/config.py), this triggers `optax.clip_by_global_norm` in the optimizer pipeline. You can also enable update-level RMS clipping using the `update_rms_clipping` parameter for finer control over update magnitudes.

### What is the difference between weight_decay_modules and default_weight_decay_mask?

`weight_decay_modules` accepts a regex string or list of patterns to explicitly match parameter names that should receive weight decay, while `default_weight_decay_mask` is a boolean flag that enables a built-in heuristic mask when no explicit patterns are provided. The `build_weight_decay_mask` function (lines 95-101) handles the logic for applying these configurations during optimizer construction.

### How does the learning rate schedule integrate with AdamConfig?

The `lr_schedule` field in `OptimizerConfig` (lines 154-179) accepts either a schedule instance or string name. When you call `build(num_train_steps)`, the configuration invokes `lr_scheduler(num_train_steps)` to compile the schedule, which is then injected as the final scaling step in the Optax optimizer chain. You can use built-in schedules like `CosineLrSchedule` or implement custom schedules that conform to the schedule protocol.

### Can I use AdamC corrected weight decay with standard Adam in Levanter?

Yes. Set `adamc_weight_decay=True` in your `AdamConfig` instantiation. This flag modifies the weight decay logic to maintain constant weight decay strength relative to the learning rate, preventing the decay magnitude from changing as learning rate schedules decay during training. The implementation applies this correction within the `build()` method's weight decay stage.