vanillatransformer_losses
vanillatransformer_losses
Pluggable point-loss functions for the VanillaTransformer forecaster. Self-contained (no dependency on other models' loss modules). Each loss has signature (pred, target) -> scalar and reduces by mean over all elements.
mae(pred, target)
vanillatransformer_losses.mae
Mean absolute error: mean(|pred - target|).
| Parameter | Type | Default | Description |
|---|---|---|---|
pred |
jnp.ndarray |
- | (undocumented) |
target |
jnp.ndarray |
- | (undocumented) |
Returns: jnp.ndarray
mse(pred, target)
vanillatransformer_losses.mse
Mean squared error: mean((pred - target)**2).
| Parameter | Type | Default | Description |
|---|---|---|---|
pred |
jnp.ndarray |
- | (undocumented) |
target |
jnp.ndarray |
- | (undocumented) |
Returns: jnp.ndarray
huber(pred, target)
vanillatransformer_losses.huber
Huber loss with delta = 1.0. Quadratic for |r| <= 1, linear outside.
| Parameter | Type | Default | Description |
|---|---|---|---|
pred |
jnp.ndarray |
- | (undocumented) |
target |
jnp.ndarray |
- | (undocumented) |
Returns: jnp.ndarray
resolve(loss)
vanillatransformer_losses.resolve
Return a callable loss from either a registry string or a callable.
| Parameter | Type | Default | Description |
|---|---|---|---|
loss |
str | LossFn |
- | (undocumented) |
Returns: LossFn
Raises: ValueError
LOSSES
vanillatransformer_losses.LOSSES
Mapping of available loss functions by name.
Type: Mapping[str, LossFn]
Value: {"mae": mae, "mse": mse, "huber": huber}