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