Skip to content

PytorchForecastingDeepAR

PytorchForecastingDeepAR

class PytorchForecastingDeepAR(model_params: dict[str, Any] | None = None, allowed_encoder_known_variable_names: list[str] | None = None, dataset_params: dict[str, Any] | None = None, train_to_dataloader_params: dict[str, Any] | None = None, validation_to_dataloader_params: dict[str, Any] | None = None, trainer_params: dict[str, Any] | None = None, model_path: str | None = None, deterministic: bool = False, random_log_path: bool = False, broadcasting: bool = False)[source]

pytorch-forecasting DeepAR model.

Parameters:
model_paramsdict[str, Any] (default=None)

parameters to be passed to initialize the pytorch-forecasting NBeats model [1] for example: {“cell_type”: “GRU”, “rnn_layers”: 3}

dataset_paramsdict[str, Any] (default=None)

parameters to initialize TimeSeriesDataSet [2] from pandas.DataFrame max_prediction_length will be overwrite according to fh time_idx, target, group_ids, time_varying_known_reals, time_varying_unknown_reals will be inferred from data, so you do not have to pass them

train_to_dataloader_paramsdict[str, Any] (default=None)

parameters to be passed for TimeSeriesDataSet.to_dataloader() by default {“train”: True}

validation_to_dataloader_paramsdict[str, Any] (default=None)

parameters to be passed for TimeSeriesDataSet.to_dataloader() by default {“train”: False}

model_path: string (default=None)

try to load a existing model without fitting. Calling the fit function is still needed, but no real fitting will be performed.

deterministic: bool (default=False)

set seed before predict, so that it will give the same output for the same input

random_log_path: bool (default=False)

use random root directory for logging. This parameter is for CI test in Github action, not designed for end users.

Attributes:
algorithm_class

Import underlying pytorch-forecasting algorithm class.

algorithm_parameters

Get keyword parameters for the DeepAR class.

dict

keyword arguments for the underlying algorithm class

cutoff

Cut-off = “present time” state of forecaster.

fh

Forecasting horizon that was passed.

is_fitted

Whether fit has been called.

state

State of the estimator.

References

Examples

>>> # import packages
>>> from sktime.forecasting.base import ForecastingHorizon
>>> from sktime.forecasting.pytorchforecasting import PytorchForecastingDeepAR
>>> from sktime.utils._testing.hierarchical import _make_hierarchical
>>> from sklearn.model_selection import train_test_split
>>> # generate random data
>>> data = _make_hierarchical(
...     hierarchy_levels=(5, 200), max_timepoints=50, min_timepoints=50, n_columns=3
... )
>>> # define forecast horizon
>>> max_prediction_length = 5
>>> fh = ForecastingHorizon(range(1, max_prediction_length + 1), is_relative=True)
>>> # split X, y data for train and test
>>> x = data[["c0", "c1"]]
>>> y = data["c2"].to_frame()
>>> X_train, X_test, y_train, y_test = train_test_split(
...     x, y, test_size=0.2, train_size=0.8, shuffle=False
... )
>>> len_levels = len(y_test.index.names)
>>> y_test = y_test.groupby(level=list(range(len_levels - 1))).apply(
...     lambda x: x.droplevel(list(range(len_levels - 1))).iloc[:-max_prediction_length]
... )
>>> # define the model
>>> model = PytorchForecastingDeepAR(
...     trainer_params={
...         "max_epochs": 5,  # for quick test
...         "limit_train_batches": 10,  # for quick test
...     },
... )
>>> # fit and predict
>>> model.fit(y=y_train, X=X_train, fh=fh)
PytorchForecastingDeepAR(trainer_params={'limit_train_batches': 10,
                                        'max_epochs': 5})
>>> y_pred = model.predict(fh, X=X_test, y=y_test)
>>> print(y_test)
                            c2
h0   h1     time
h0_0 h1_180 2000-01-01  5.006716
            2000-01-02  5.197903
            2000-01-03  4.477552
            2000-01-04  4.751521
            2000-01-05  3.323994
...                          ...
h0_4 h1_199 2000-02-10  5.590399
            2000-02-11  5.595445
            2000-02-12  4.915307
            2000-02-13  4.726925
            2000-02-14  5.482842

[4500 rows x 1 columns] >>> print(y_pred) # doctest: +SKIP

c2

h0 h1 time h0_0 h1_180 2000-02-15 4.919366

2000-02-16 4.862666 2000-02-17 5.021425 2000-02-18 4.934844 2000-02-19 4.808967

… … h0_4 h1_199 2000-02-15 5.150748

2000-02-16 5.230827 2000-02-17 5.123736 2000-02-18 5.139505 2000-02-19 5.121511

[500 rows x 1 columns]

Methods

check_is_fitted([method_name])

Check if the estimator has been fitted.

clone()

Obtain a clone of the object with same hyper-parameters and config.

clone_tags(estimator[, tag_names])

Clone tags from another object as dynamic override.

create_test_instance([parameter_set])

Construct an instance of the class, using first test parameter set.

create_test_instances_and_names([parameter_set])

Create list of all test instances and a list of names for them.

fit(y[, X, fh])

Fit forecaster to training data.

fit_predict(y[, X, fh, X_pred])

Fit and forecast time series at future horizon.

get_class_tag(tag_name[, tag_value_default])

Get class tag value from class, with tag level inheritance from parents.

get_class_tags()

Get class tags from class, with tag level inheritance from parent classes.

get_config()

Get config flags for self.

get_fitted_params([deep])

Get fitted parameters.

get_param_defaults()

Get object's parameter defaults.

get_param_names([sort])

Get object's parameter names.

get_params([deep])

Get a dict of parameters values for this object.

get_pretrained_params([deep])

Get pretrained parameters of this estimator.

get_tag(tag_name[, tag_value_default, ...])

Get tag value from instance, with tag level inheritance and overrides.

get_tags()

Get tags from instance, with tag level inheritance and overrides.

get_test_params([parameter_set])

Return testing parameter settings for the estimator.

is_composite()

Check if the object is composed of other BaseObjects.

load_from_path(serial)

Load object from file location.

load_from_serial(serial)

Load object from serialized memory container.

predict([fh, X])

Forecast time series at future horizon.

predict_interval([fh, X, coverage])

Compute/return prediction interval forecasts.

predict_proba([fh, X, marginal])

Compute/return fully probabilistic forecasts.

predict_quantiles([fh, X, alpha])

Compute/return quantile forecasts.

predict_residuals([y, X])

Return residuals of time series forecasts.

predict_var([fh, X, cov])

Compute/return variance forecasts.

pretrain(y[, X, fh])

Pre-train forecaster on panel (global) data.

reset()

Reset the object to a clean post-init state.

save([path, serialization_format])

Save serialized self to bytes-like object or to (.zip) file.

score(y[, X, fh])

Scores forecast against ground truth, using MAPE (non-symmetric).

set_config(**config_dict)

Set config flags to given values.

set_params(**params)

Set the parameters of this object.

set_random_state([random_state, deep, ...])

Set random_state pseudo-random seed parameters for self.

set_tags(**tag_dict)

Set instance level tag overrides to given values.

update(y[, X, update_params])

Update cutoff value and, optionally, fitted parameters.

update_predict(y[, cv, X, update_params, ...])

Make predictions and update model iteratively over the test set.

update_predict_single([y, fh, X, update_params])

Update model with new data and make forecasts.