LossMixin#
- class liesel.optim.LossMixin[source]#
Bases:
objectShared convenience implementation for differentiable losses.
Subclasses must define
splitandloss_train_batched(). They should also defineloss_train()if they support exact full-data monitoring. The mixin provides validation-position helpers and JAX gradient methods used byOptimizer.Examples
A minimal quadratic loss can inherit from
LossMixinand immediately use the gradient helpers:>>> import jax.numpy as jnp >>> from liesel.optim import LossMixin, PositionSplit >>> from liesel.optim.types import Position >>> class Quadratic(LossMixin): ... def __init__(self): ... self.split = PositionSplit( ... Position({"y": jnp.array([0.0])}), ... Position({}), ... Position({}), ... 1, ... 0, ... 0, ... ) ... ... def loss_train_batched(self, params, carry): ... del carry ... return params["x"] ** 2 >>> loss = Quadratic() >>> loss.grad(Position({"x": jnp.array(3.0)}), carry=None)["x"] Array(6., dtype=float32, weak_type=True) >>> loss.obs_validate["y"].tolist() [0.0]
Methods
grad(params, carry)Computes the gradient of
loss_train_batched().loss_train(params, carry)Computes the full-data training loss.
value_and_grad(params, carry)Evaluates
loss_train_batched()and its gradient.Attributes
Observed position used for validation.
Number of observations used by validation.
Scalar validation likelihood scale.
Train/validation/test split used by the loss.
Training objective differentiated by
grad()andvalue_and_grad().