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.

Loss functions for Chronax TiDE.

Mirrors the RNN / N-BEATS loss API: masked_mae(y, y_hat, mask=None) -> scalar masked_mse(y, y_hat, mask=None) -> scalar Shapes follow the batch convention: y and y_hat are [B, h, output_size] (or any broadcast-compatible shape). The optional mask is [B, h, 1].

masked_mae

chronax.loss.masked_mae

Mean Absolute Error, optionally masked.

masked_mae(y, y_hat, mask=None)

Parameter Type Default Description
y jnp.ndarray - Target values [B, h, output_size] (or broadcastable).
y_hat jnp.ndarray - Predicted values, same shape as y.
mask Optional[jnp.ndarray] None Float mask broadcastable to y; 1 = include, 0 = exclude.

Returns: jnp.ndarray (Scalar loss value.)

masked_mse

chronax.loss.masked_mse

Mean Squared Error, optionally masked.

masked_mse(y, y_hat, mask=None)

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

Returns: jnp.ndarray

resolve

chronax.loss.resolve

Return a callable loss from either a registry string or a callable.

resolve(loss)

Parameter Type Default Description
loss - - (undocumented)

Returns: Callable