How to Configure AdamConfig and Optimizer Settings in Levanter
Levanter's AdamConfig class in 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 (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.
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.
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.
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, which implements logic to detect and skip optimizer steps when gradients explode beyond a threshold.
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:
- Optional global-norm clipping via
optax.clip_by_global_normwhenmax_grad_normis set - Adam core update rule via
optax.scale_by_adamusingbeta1,beta2, andepsilon - Optional weight decay via
optax.add_decayed_weightsusing the mask built fromweight_decay_modules - Optional RMS-clipping or custom update-norm clipping
- Learning rate scaling using the injected schedule
- Optional step-skipping wrapper when
skip_bad_stepsis 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 and 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, line 37).
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
AdamConfiglives inlib/levanter/src/levanter/optim/config.py(lines 57-71) and inherits fromOptimizerConfig- Configure learning rate schedules through the
lr_schedulefield accepting objects likeCosineLrScheduleor string identifiers - Control weight decay application using
weight_decay_modulesregex patterns ordefault_weight_decay_mask, with AdamC correction available viaadamc_weight_decay - Apply gradient clipping through
max_grad_normandupdate_rms_clippingparameters - Build functional optimizers by calling
.build(num_train_steps), which returns anoptax.GradientTransformation - Integrate into training scripts using dataclass fields with
default_factory=AdamConfigto 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, 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.
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 →