TrainState
train.TrainState · inherits flax.training.train_state.TrainState
Standard Flax TrainState; aliased here for forward compatibility.
__init__(self, ...)
(Inherited from flax.training.train_state.TrainState)
make_lr_schedule(learning_rate, max_steps, num_lr_decays)
StepLR-style schedule that halves the LR num_lr_decays times.
Returns a constant LR when num_lr_decays <= 0. Otherwise the LR is piecewise-constant, multiplied by 0.5 at evenly spaced step boundaries -- the same gamma=0.5 StepLR recipe neuralforecast uses.
| Parameter | Type | Default | Description |
|---|---|---|---|
| learning_rate | float | - | (undocumented) |
| max_steps | int | - | (undocumented) |
| num_lr_decays | int | - | (undocumented) |
Returns: optax.ScalarOrSchedule
sample_batch_indices(key, n_train, batch_size, n_steps)
Sample [n_steps, batch_size] window indices (NF-style).
With replacement when n_train < batch_size (NF torch.randint); without replacement via permutation slice otherwise (NF randperm[:B]).
| Parameter | Type | Default | Description |
|---|---|---|---|
| key | jax.Array | - | (undocumented) |
| n_train | int | - | (undocumented) |
| batch_size | int | - | (undocumented) |
| n_steps | int | - | (undocumented) |
Returns: jax.Array
should_stop_early(*, early_stop_patience_steps, checks_without_improvement)
True when validation has not improved for patience consecutive checks.
| Parameter | Type | Default | Description |
|---|---|---|---|
| early_stop_patience_steps | int | - | (undocumented) |
| checks_without_improvement | int | - | (undocumented) |
Returns: bool
create_train_state(rng, config, learning_rate=1e-4, weight_decay=0.0, grad_clip=1.0, num_lr_decays=-1, max_steps=1000, optimizer=None)
Initialise model parameters and optimiser state.
| Parameter | Type | Default | Description |
|---|---|---|---|
| rng | jax.Array | - | PRNG key for parameter initialisation. |
| config | FEDformerConfig | - | Model architecture config. |
| learning_rate | float | 1e-4 | Peak learning rate. |
| weight_decay | float | 0.0 | If > 0 use optax.adamw, else optax.adam. |
| grad_clip | float | 1.0 | Global gradient-norm clip threshold (0 = disabled). |
| num_lr_decays | int | -1 | Number of StepLR decays (gamma=0.5); -1 = constant LR. |
| max_steps | int | 1000 | Total training steps (sets decay-schedule spacing). |
| optimizer | Optional[optax.GradientTransformation] | None | Optional pre-built transform; overrides all defaults. |
Returns: TrainState
train_step(state, batch, rng, loss_fn=masked_mae)
JIT-compiled training step on a pre-scaled batch.
| Parameter | Type | Default | Description |
|---|---|---|---|
| state | TrainState | - | (undocumented) |
| batch | Dict[str, Optional[jnp.ndarray]] | - | (undocumented) |
| rng | jax.Array | - | (undocumented) |
| loss_fn | Callable | masked_mae | (undocumented) |
Returns: Tuple[TrainState, jnp.ndarray, jnp.ndarray] (Returns (new_state, loss, predictions).)
eval_step(state, batch, loss_fn=masked_mae)
JIT-compiled evaluation step. Deterministic, no gradients.
| Parameter | Type | Default | Description |
|---|---|---|---|
| state | TrainState | - | (undocumented) |
| batch | Dict[str, Optional[jnp.ndarray]] | - | (undocumented) |
| loss_fn | Callable | masked_mae | (undocumented) |
Returns: Tuple[jnp.ndarray, jnp.ndarray]
train_window_step(state, windows, masks, rng, input_size, loss_fn=_mae)
JIT-compiled training step on raw [B, input_size + h] windows.
Per-window RobustScaler is applied internally. input_size is static (changing it triggers a recompile).
| Parameter | Type | Default | Description |
|---|---|---|---|
| state | TrainState | - | (undocumented) |
| windows | jnp.ndarray | - | (undocumented) |
| masks | jnp.ndarray | - | (undocumented) |
| rng | jax.Array | - | (undocumented) |
| input_size | int | - | (undocumented) |
| loss_fn | Callable | _mae | (undocumented) |
Returns: Tuple[TrainState, jnp.ndarray, jnp.ndarray]
scan_train_steps(state, all_windows, all_masks, batch_idx, rng_seq, input_size, loss_fn=_mae)
Run lax.scan over precomputed window indices (no per-step host sync).
| Parameter | Type | Default | Description |
|---|---|---|---|
| state | TrainState | - | (undocumented) |
| all_windows | jnp.ndarray | - | Full train windows [n_train, L+h]. |
| all_masks | jnp.ndarray | - | Matching availability masks. |
| batch_idx | jnp.ndarray | - | [n_steps, batch_size] indices into all_windows. |
| rng_seq | jnp.ndarray | - | [n_steps, 2] dropout keys. |
| input_size | int | - | (undocumented) |
| loss_fn | Callable | _mae | (undocumented) |
Returns: Tuple[TrainState, jnp.ndarray]
eval_window_step(state, windows, masks, input_size, loss_fn=_mae)
JIT-compiled evaluation on raw windows. Deterministic, no gradients.
| Parameter | Type | Default | Description |
|---|---|---|---|
| state | TrainState | - | (undocumented) |
| windows | jnp.ndarray | - | (undocumented) |
| masks | jnp.ndarray | - | (undocumented) |
| input_size | int | - | (undocumented) |
| loss_fn | Callable | _mae | (undocumented) |
Returns: Tuple[jnp.ndarray, jnp.ndarray]
train_loop(state, train_batches, *, num_epochs=10, rng=None, eval_batches=None, loss_fn=masked_mae)
Drive train_step / eval_step over multiple epochs.
train_batches may be a list (re-iterated each epoch) or a generator (materialise it first if you need multiple epochs).
| Parameter | Type | Default | Description |
|---|---|---|---|
| state | TrainState | - | (undocumented) |
| train_batches | Iterable[Dict[str, Optional[jnp.ndarray]]] | - | (undocumented) |
| num_epochs | int | 10 | (undocumented) |
| rng | Optional[jax.Array] | None | (undocumented) |
| eval_batches | Optional[Iterable[Dict[str, Optional[jnp.ndarray]]]] | None | (undocumented) |
| loss_fn | Callable | masked_mae | (undocumented) |
Returns: Tuple[TrainState, List[Dict[str, float]]]
predict_step(state, context, *, h, input_size, scaler=None)
Deterministic forecast from a 1-D context window of length input_size.
| Parameter | Type | Default | Description |
|---|---|---|---|
| state | TrainState | - | (undocumented) |
| context | jnp.ndarray | - | (undocumented) |
| h | int | - | (undocumented) |
| input_size | int | - | (undocumented) |
| scaler | Optional[RobustScaler] | None | (undocumented) |
Returns: jnp.ndarray (Returns a 1-D array of shape (h,) on the original scale.)
train(y, *, config, max_steps=1000, learning_rate=1e-4, batch_size=32, num_lr_decays=3, val_fraction=0.1, val_check_steps=100, early_stop_patience_steps=-1, grad_clip=1.0, weight_decay=0.0, loss_fn=_mae, random_seed=1, verbose=False)
Train on a univariate 1-D series; return best-validation TrainState.
| Parameter | Type | Default | Description |
|---|---|---|---|
| y | jnp.ndarray | - | (undocumented) |
| config | FEDformerConfig | - | (undocumented) |
| max_steps | int | 1000 | (undocumented) |
| learning_rate | float | 1e-4 | (undocumented) |
| batch_size | int | 32 | (undocumented) |
| num_lr_decays | int | 3 | (undocumented) |
| val_fraction | float | 0.1 | (undocumented) |
| val_check_steps | int | 100 | (undocumented) |
| early_stop_patience_steps | int | -1 | (undocumented) |
| grad_clip | float | 1.0 | (undocumented) |
| weight_decay | float | 0.0 | (undocumented) |
| loss_fn | Callable | _mae | (undocumented) |
| random_seed | int | 1 | (undocumented) |
| verbose | bool | False | (undocumented) |
Returns: TrainState
Raises: RuntimeError if a non-finite training loss is observed.