From eeeb96335a1b8e0a79a931e8d3ad71a41c402ebc Mon Sep 17 00:00:00 2001 From: "Dr. Kashif Rasul" Date: Sat, 2 Jan 2021 10:59:50 +0100 Subject: [PATCH] updated tempflow --- pts/model/deepar/deepar_estimator.py | 1 + pts/model/deepvar/deepvar_estimator.py | 2 ++ pts/model/lstnet/lstnet_estimator.py | 3 +- pts/model/n_beats/n_beats_estimator.py | 2 ++ .../simple_feedforward_estimator.py | 1 + pts/model/tempflow/tempflow_estimator.py | 34 ++++++++++++------- pts/model/tempflow/tempflow_network.py | 1 + .../transformer/transformer_estimator.py | 2 ++ .../transformer_tempflow_estimator.py | 33 +++++++++++------- .../transformer_tempflow_network.py | 3 +- 10 files changed, 54 insertions(+), 28 deletions(-) diff --git a/pts/model/deepar/deepar_estimator.py b/pts/model/deepar/deepar_estimator.py index 1924498..eb10f92 100644 --- a/pts/model/deepar/deepar_estimator.py +++ b/pts/model/deepar/deepar_estimator.py @@ -27,6 +27,7 @@ from gluonts.torch.support.util import copy_parameters from gluonts.torch.model.predictor import PyTorchPredictor from gluonts.torch.modules.distribution_output import DistributionOutput from gluonts.model.predictor import Predictor + from pts.model.utils import get_module_forward_input_names from pts import Trainer from pts.model import PyTorchEstimator diff --git a/pts/model/deepvar/deepvar_estimator.py b/pts/model/deepvar/deepvar_estimator.py index 30429f4..cf6f8c2 100644 --- a/pts/model/deepvar/deepvar_estimator.py +++ b/pts/model/deepvar/deepvar_estimator.py @@ -27,6 +27,7 @@ from gluonts.transform import ( AddAgeFeature, cdf_to_gaussian_forward_transform, ) + from pts import Trainer from pts.model.utils import get_module_forward_input_names from pts.feature import ( @@ -35,6 +36,7 @@ from pts.feature import ( ) from pts.model import PyTorchEstimator from pts.modules import LowRankMultivariateNormalOutput + from .deepvar_network import DeepVARTrainingNetwork, DeepVARPredictionNetwork diff --git a/pts/model/lstnet/lstnet_estimator.py b/pts/model/lstnet/lstnet_estimator.py index 9f8baf2..3b448e8 100644 --- a/pts/model/lstnet/lstnet_estimator.py +++ b/pts/model/lstnet/lstnet_estimator.py @@ -6,7 +6,6 @@ import torch.nn as nn from gluonts.dataset.field_names import FieldName from gluonts.torch.support.util import copy_parameters -from pts.model.utils import get_module_forward_input_names from gluonts.torch.model.predictor import PyTorchPredictor from gluonts.model.predictor import Predictor from gluonts.transform import ( @@ -17,8 +16,10 @@ from gluonts.transform import ( AddObservedValuesIndicator, AsNumpyArray, ) + from pts.model import PyTorchEstimator from pts import Trainer +from pts.model.utils import get_module_forward_input_names from .lstnet_network import LSTNetTrain, LSTNetPredict diff --git a/pts/model/n_beats/n_beats_estimator.py b/pts/model/n_beats/n_beats_estimator.py index 221f486..fc867a3 100644 --- a/pts/model/n_beats/n_beats_estimator.py +++ b/pts/model/n_beats/n_beats_estimator.py @@ -14,9 +14,11 @@ from gluonts.transform import ( 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, diff --git a/pts/model/simple_feedforward/simple_feedforward_estimator.py b/pts/model/simple_feedforward/simple_feedforward_estimator.py index 877bf71..12a21c2 100644 --- a/pts/model/simple_feedforward/simple_feedforward_estimator.py +++ b/pts/model/simple_feedforward/simple_feedforward_estimator.py @@ -19,6 +19,7 @@ from gluonts.transform import ( InstanceSplitter, ExpectedNumInstanceSampler, ) + from pts.model.utils import get_module_forward_input_names from pts import Trainer from pts.model import PyTorchEstimator diff --git a/pts/model/tempflow/tempflow_estimator.py b/pts/model/tempflow/tempflow_estimator.py index 0e6e4e7..1003ae9 100644 --- a/pts/model/tempflow/tempflow_estimator.py +++ b/pts/model/tempflow/tempflow_estimator.py @@ -2,15 +2,13 @@ from typing import List, Optional import torch -from pts import Trainer -from pts.dataset import FieldName -from pts.feature import ( - TimeFeature, - fourier_time_features_from_frequency_str, - get_fourier_lags_for_frequency, -) -from pts.model import PyTorchEstimator, PyTorchPredictor, copy_parameters -from pts.transform import ( +from gluonts.dataset.field_names import FieldName +from gluonts.time_feature import TimeFeature +from gluonts.torch.model.predictor import PyTorchPredictor +from gluonts.torch.support.util import copy_parameters +from gluonts.model.predictor import Predictor +from gluonts.torch.model.predictor import PyTorchPredictor +from gluonts.transform import ( Transformation, Chain, InstanceSplitter, @@ -24,6 +22,15 @@ from pts.transform import ( SetFieldIfNotPresent, TargetDimIndicator, ) + +from pts import Trainer +from pts.feature import ( + fourier_time_features_from_frequency, + lags_for_fourier_time_features_from_frequency +) +from pts.model.utils import get_module_forward_input_names +from pts.model import PyTorchEstimator + from .tempflow_network import TempFlowTrainingNetwork, TempFlowPredictionNetwork @@ -83,13 +90,13 @@ class TempFlowEstimator(PyTorchEstimator): self.lags_seq = ( lags_seq if lags_seq is not None - else get_fourier_lags_for_frequency(freq_str=freq) + else lags_for_fourier_time_features_from_frequency(freq_str=freq) ) self.time_features = ( time_features if time_features is not None - else fourier_time_features_from_frequency_str(self.freq) + else fourier_time_features_from_frequency(self.freq) ) self.history_length = self.context_length + max(self.lags_seq) @@ -181,7 +188,7 @@ class TempFlowEstimator(PyTorchEstimator): transformation: Transformation, trained_network: TempFlowTrainingNetwork, device: torch.device, - ) -> PyTorchPredictor: + ) -> Predictor: prediction_network = TempFlowPredictionNetwork( input_size=self.input_size, target_dim=self.target_dim, @@ -206,13 +213,14 @@ class TempFlowEstimator(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, prediction_length=self.prediction_length, device=device, - output_transform=None, ) diff --git a/pts/model/tempflow/tempflow_network.py b/pts/model/tempflow/tempflow_network.py index f954dfa..028efdf 100644 --- a/pts/model/tempflow/tempflow_network.py +++ b/pts/model/tempflow/tempflow_network.py @@ -4,6 +4,7 @@ import torch import torch.nn as nn from gluonts.core.component import validated + from pts.model import weighted_average from pts.modules import RealNVP, MAF, FlowOutput, MeanScaler, NOPScaler diff --git a/pts/model/transformer/transformer_estimator.py b/pts/model/transformer/transformer_estimator.py index bce6639..b3227e4 100644 --- a/pts/model/transformer/transformer_estimator.py +++ b/pts/model/transformer/transformer_estimator.py @@ -23,6 +23,7 @@ from gluonts.transform import ( VstackFeatures, SetField, ) + from pts import Trainer from pts.model.utils import get_module_forward_input_names from pts.feature import ( @@ -31,6 +32,7 @@ from pts.feature import ( ) from pts.model import PyTorchEstimator from pts.modules import StudentTOutput + from .transformer_network import ( TransformerTrainingNetwork, TransformerPredictionNetwork, diff --git a/pts/model/transformer_tempflow/transformer_tempflow_estimator.py b/pts/model/transformer_tempflow/transformer_tempflow_estimator.py index d489d35..adcb092 100644 --- a/pts/model/transformer_tempflow/transformer_tempflow_estimator.py +++ b/pts/model/transformer_tempflow/transformer_tempflow_estimator.py @@ -2,15 +2,12 @@ from typing import List, Optional import torch -from pts import Trainer -from pts.dataset import FieldName -from pts.feature import ( - TimeFeature, - fourier_time_features_from_frequency_str, - get_fourier_lags_for_frequency, -) -from pts.model import PyTorchEstimator, PyTorchPredictor, copy_parameters -from pts.transform import ( +from gluonts.dataset.field_names import FieldName +from gluonts.time_feature import TimeFeature +from gluonts.torch.support.util import copy_parameters +from gluonts.torch.model.predictor import PyTorchPredictor +from gluonts.model.predictor import Predictor +from gluonts.transform import ( Transformation, Chain, InstanceSplitter, @@ -24,6 +21,15 @@ from pts.transform import ( SetFieldIfNotPresent, TargetDimIndicator, ) + +from pts import Trainer +from pts.model import PyTorchEstimator +from pts.model.utils import get_module_forward_input_names +from pts.feature import ( + fourier_time_features_from_frequency, + lags_for_fourier_time_features_from_frequency, +) + from .transformer_tempflow_network import ( TransformerTempFlowTrainingNetwork, TransformerTempFlowPredictionNetwork, @@ -94,13 +100,13 @@ class TransformerTempFlowEstimator(PyTorchEstimator): self.lags_seq = ( lags_seq if lags_seq is not None - else get_fourier_lags_for_frequency(freq_str=freq) + else lags_for_fourier_time_features_from_frequency(freq_str=freq) ) self.time_features = ( time_features if time_features is not None - else fourier_time_features_from_frequency_str(self.freq) + else fourier_time_features_from_frequency(self.freq) ) self.history_length = self.context_length + max(self.lags_seq) @@ -197,7 +203,7 @@ class TransformerTempFlowEstimator(PyTorchEstimator): transformation: Transformation, trained_network: TransformerTempFlowTrainingNetwork, device: torch.device, - ) -> PyTorchPredictor: + ) -> Predictor: prediction_network = TransformerTempFlowPredictionNetwork( input_size=self.input_size, target_dim=self.target_dim, @@ -225,13 +231,14 @@ class TransformerTempFlowEstimator(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, prediction_length=self.prediction_length, device=device, - output_transform=None, ) diff --git a/pts/model/transformer_tempflow/transformer_tempflow_network.py b/pts/model/transformer_tempflow/transformer_tempflow_network.py index 69e5768..e1a09a7 100644 --- a/pts/model/transformer_tempflow/transformer_tempflow_network.py +++ b/pts/model/transformer_tempflow/transformer_tempflow_network.py @@ -3,7 +3,8 @@ from typing import List, Optional, Tuple import torch import torch.nn as nn -from gluonts.core.component import validated +from gluonts.core.component import + from pts.modules import RealNVP, MAF, FlowOutput, MeanScaler, NOPScaler