Esc
Ask AIAnswers may be inaccurate; check the linked pages.Esc
Ask anything about these docs, like how to get started or what a function does.

mae

loss.py

Mean absolute error; masked mean when mask is given (NF padder).

mae(pred, target, mask=None)

Parameter Type Default Description
pred jnp.ndarray - (undocumented)
target jnp.ndarray - (undocumented)
mask Optional[jnp.ndarray] None (undocumented)

Returns: jnp.ndarray

mse

loss.py

Mean squared error; masked mean when mask is given.

mse(pred, target, mask=None)

Parameter Type Default Description
pred jnp.ndarray - (undocumented)
target jnp.ndarray - (undocumented)
mask Optional[jnp.ndarray] None (undocumented)

Returns: jnp.ndarray

huber

loss.py

Huber loss with delta = 1.0 (smooth L1); masked mean when mask is given.

Quadratic for |r| <= 1, linear outside. Smooth at the transition so gradients are well-behaved everywhere.

huber(pred, target, mask=None)

Parameter Type Default Description
pred jnp.ndarray - (undocumented)
target jnp.ndarray - (undocumented)
mask Optional[jnp.ndarray] None (undocumented)

Returns: jnp.ndarray

LOSSES

loss.py

A mapping of loss name strings to callable loss functions ({"mae": mae, "mse": mse, "huber": huber}).

resolve

loss.py

Return a callable from a registry key or pass-through a callable.

resolve(loss)

Parameter Type Default Description
loss str \| LossFn - (undocumented)

Returns: LossFn

masked_mae

loss.py

Mean Absolute Error with multiplicative masking.

masked_mae(y, y_hat, mask=None, horizon_weight=None)

Parameter Type Default Description
y jnp.ndarray - Target values [B, H, 1].
y_hat jnp.ndarray - Predictions [B, H, 1].
mask Optional[jnp.ndarray] None Optional [B, H, 1] mask (1 = include, 0 = ignore).
horizon_weight Optional[jnp.ndarray] None Optional [H] per-step weights.

Returns: jnp.ndarray

masked_mse

loss.py

Mean Squared Error with multiplicative masking.

masked_mse(y, y_hat, mask=None, horizon_weight=None)

Parameter Type Default Description
y jnp.ndarray - (undocumented)
y_hat jnp.ndarray - (undocumented)
mask Optional[jnp.ndarray] None (undocumented)
horizon_weight Optional[jnp.ndarray] None (undocumented)

Returns: jnp.ndarray