LossMixin

Contents

LossMixin#

class liesel.optim.LossMixin[source]#

Bases: object

Shared convenience implementation for differentiable losses.

Subclasses must define split and loss_train_batched(). They should also define loss_train() if they support exact full-data monitoring. The mixin provides validation-position helpers and JAX gradient methods used by Optimizer.

Examples

A minimal quadratic loss can inherit from LossMixin and 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

obs_validate

Observed position used for validation.

validate_axis_size

Number of observations used by validation.

validate_sample_scale

Scalar validation likelihood scale.

split

Train/validation/test split used by the loss.

loss_train_batched

Training objective differentiated by grad() and value_and_grad().