From 4fc70dc7b5a7c372c6b60a7612efd016612a01a7 Mon Sep 17 00:00:00 2001 From: "Dr. Kashif Rasul" Date: Fri, 1 Jan 2021 22:33:40 +0100 Subject: [PATCH] updated nbeats --- pts/model/n_beats/n_beats_ensemble.py | 10 +++++++--- pts/model/n_beats/n_beats_estimator.py | 14 ++++++++++---- pts/model/n_beats/n_beats_network.py | 2 +- 3 files changed, 18 insertions(+), 8 deletions(-) diff --git a/pts/model/n_beats/n_beats_ensemble.py b/pts/model/n_beats/n_beats_ensemble.py index 3c0b86b..b11626f 100644 --- a/pts/model/n_beats/n_beats_ensemble.py +++ b/pts/model/n_beats/n_beats_ensemble.py @@ -5,9 +5,13 @@ from typing import List, Optional, Iterator import numpy as np +from gluonts.dataset.field_names import FieldName +from gluonts.dataset.common import Dataset +from gluonts.model.predictor import Predictor +from gluonts.model.forecast import Forecast, SampleForecast + +from pts.model import PyTorchEstimator from pts import Trainer -from pts.dataset import Dataset, FieldName -from pts.model import Predictor, SampleForecast, Forecast, Estimator from .n_beats_estimator import NBEATSEstimator from .n_beats_network import VALID_LOSS_FUNCTIONS @@ -89,7 +93,7 @@ class NBEATSEnsemblePredictor(Predictor): ) -class NBEATSEnsembleEstimator(Estimator): +class NBEATSEnsembleEstimator(PyTorchEstimator): """ An ensemble N-BEATS Estimator (approximately) as described in the paper: https://arxiv.org/abs/1905.10437. diff --git a/pts/model/n_beats/n_beats_estimator.py b/pts/model/n_beats/n_beats_estimator.py index 272e69b..221f486 100644 --- a/pts/model/n_beats/n_beats_estimator.py +++ b/pts/model/n_beats/n_beats_estimator.py @@ -3,16 +3,20 @@ from typing import List, Optional import torch import torch.nn as nn -from pts import Trainer -from pts.dataset import FieldName -from pts.model import PyTorchEstimator, Predictor, PyTorchPredictor, copy_parameters -from pts.transform import ( +from gluonts.dataset.field_names import FieldName +from gluonts.model.predictor import Predictor +from gluonts.torch.model.predictor import PyTorchPredictor +from gluonts.torch.support.util import copy_parameters +from gluonts.transform import ( InstanceSplitter, Transformation, Chain, RemoveFields, ExpectedNumInstanceSampler, ) +from pts import Trainer +from pts.model import PyTorchEstimator +from pts.model.utils import get_module_forward_input_names from .n_beats_network import ( NBEATSPredictionNetwork, NBEATSTrainingNetwork, @@ -179,9 +183,11 @@ class NBEATSEstimator(PyTorchEstimator): ).to(device) copy_parameters(trained_network, prediction_network) + input_names = get_module_forward_input_names(prediction_network) return PyTorchPredictor( input_transform=transformation, + input_names=input_names, prediction_net=prediction_network, batch_size=self.trainer.batch_size, freq=self.freq, diff --git a/pts/model/n_beats/n_beats_network.py b/pts/model/n_beats/n_beats_network.py index 3d3aae9..394b521 100644 --- a/pts/model/n_beats/n_beats_network.py +++ b/pts/model/n_beats/n_beats_network.py @@ -5,7 +5,7 @@ import torch import torch.nn as nn import torch.nn.functional as F -from pts.feature import get_seasonality +from gluonts.time_feature import get_seasonality VALID_N_BEATS_STACK_TYPES = "G", "S", "T" VALID_LOSS_FUNCTIONS = "sMAPE", "MASE", "MAPE"