Optimizer#
- class liesel.optim.Optimizer(position_keys, optimizer, identifier='', activate_after_epochs=0)[source]#
Bases:
objectWraps an Optax gradient transformation for selected position entries.
Optimizeris the default adapter used byOptimEngine. It extracts the entries named inposition_keys, differentiates the configured loss with respect to only those entries, applies an Optax update, and merges the updated subset back into the full engine position.- Parameters:
position_keys (
Sequence[str]) – Names of the parameter entries owned by this optimizer.optimizer (
GradientTransformation) – Optax gradient transformation, for exampleoptax.adam(...)oroptax.sgd(...).identifier (
str, default:'') – Optional identifier used to store this optimizer’s state inOptimCarry. Missing identifiers are filled byOptimEngine.activate_after_epochs (
int, default:0) – Number of completed epochs required before this optimizer participates in batch updates.0activates it from the first epoch.
Notes
Multiple optimizers can be used in the same engine, but their
position_keysand identifiers must be disjoint after automatic naming.position_keysare normalized to a tuple during initialization.Examples
>>> import jax.numpy as jnp >>> import optax >>> from liesel.optim import Optimizer >>> from liesel.optim.types import Position >>> optimizer = Optimizer(["x"], optax.sgd(0.1), identifier="x_opt") >>> position = Position({"x": jnp.array(1.0), "y": jnp.array(2.0)}) >>> optimizer.position(position)["x"].tolist() 1.0 >>> sorted(optimizer.not_position(position)) ['y'] >>> repr(optimizer) "Optimizer(('x',), identifier=x_opt)"
Methods
init(position)Initializes the wrapped Optax transformation.
not_position(position)Extracts the subset of
positionnot owned by this optimizer.position(position)Extracts the subset of
positionowned by this optimizer.step(position, loss, carry)Runs one optimizer step on
position.Attributes
Number of completed epochs required before this optimizer participates in batch updates.
Optional identifier used to store this optimizer's state in
OptimCarry.Names of the parameter entries owned by this optimizer.
Optax gradient transformation, for example
optax.adam(...)oroptax.sgd(...).