Optimizer

Contents

Optimizer#

class liesel.optim.Optimizer(position_keys, optimizer, identifier='', activate_after_epochs=0)[source]#

Bases: object

Wraps an Optax gradient transformation for selected position entries.

Optimizer is the default adapter used by OptimEngine. It extracts the entries named in position_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 example optax.adam(...) or optax.sgd(...).

  • identifier (str, default: '') – Optional identifier used to store this optimizer’s state in OptimCarry. Missing identifiers are filled by OptimEngine.

  • activate_after_epochs (int, default: 0) – Number of completed epochs required before this optimizer participates in batch updates. 0 activates it from the first epoch.

Notes

Multiple optimizers can be used in the same engine, but their position_keys and identifiers must be disjoint after automatic naming. position_keys are 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 position not owned by this optimizer.

position(position)

Extracts the subset of position owned by this optimizer.

step(position, loss, carry)

Runs one optimizer step on position.

Attributes

activate_after_epochs

Number of completed epochs required before this optimizer participates in batch updates.

identifier

Optional identifier used to store this optimizer's state in OptimCarry.

position_keys

Names of the parameter entries owned by this optimizer.

optimizer

Optax gradient transformation, for example optax.adam(...) or optax.sgd(...).