diff --git a/pts/dataset/__init__.py b/pts/dataset/__init__.py index 0e8a474..3675109 100644 --- a/pts/dataset/__init__.py +++ b/pts/dataset/__init__.py @@ -9,6 +9,6 @@ from .sampler import ( UniformSplitSampler, ) from .process import ProcessStartField, ProcessDataEntry -from .utils import to_pandas +from .utils import to_pandas, from .stat import ScaleHistogram, calculate_dataset_statistics from .artificial import constant_dataset \ No newline at end of file diff --git a/pts/model/__init__.py b/pts/model/__init__.py index 66ea016..3338ed7 100644 --- a/pts/model/__init__.py +++ b/pts/model/__init__.py @@ -2,4 +2,4 @@ from .estimator import Estimator, PTSEstimator from .forecast import Forecast from .predictor import Predictor from .quantile import Quantile -from .utils import get_module_forward_input_names \ No newline at end of file +from .utils import get_module_forward_input_names, copy_parameters \ No newline at end of file diff --git a/pts/model/deepar/deepar_estimator.py b/pts/model/deepar/deepar_estimator.py index d2975b4..176c085 100644 --- a/pts/model/deepar/deepar_estimator.py +++ b/pts/model/deepar/deepar_estimator.py @@ -24,8 +24,8 @@ from pts.dataset import FieldName, ExpectedNumInstanceSampler from pts.model import PTSEstimator from pts.modules import DistributionOutput, StudentTOutput -from .deepar_network import DeepARTrainingNetwork - +from .deepar_network import DeepARTrainingNetwork, DeepARPredictionNetwork +from ..utils import copy_parameters class DeepAREstimator(PTSEstimator): def __init__( @@ -167,3 +167,27 @@ class DeepAREstimator(PTSEstimator): lags_seq=self.lags_seq, scaling=self.scaling, dtype=self.dtype).to(device) + + def create_predictor( + self, transformation: Transformation, trained_network: nn.Module + ) -> Predictor: + prediction_network = DeepARPredictionNetwork( + num_parallel_samples=self.num_parallel_samples, + input_size=self.input_size, + num_layers=self.num_layers, + num_cells=self.num_cells, + cell_type=self.cell_type, + history_length=self.history_length, + context_length=self.context_length, + prediction_length=self.prediction_length, + distr_output=self.distr_output, + dropout_rate=self.dropout_rate, + cardinality=self.cardinality, + embedding_dimension=self.embedding_dimension, + lags_seq=self.lags_seq, + scaling=self.scaling, + dtype=self.dtype).to(trained_network.device) + + copy_parameters(trained_network, prediction_network) + + return RepresentableBlockPredictor() \ No newline at end of file diff --git a/pts/model/deepar/deepar_network.py b/pts/model/deepar/deepar_network.py index e6f3729..1f726e5 100644 --- a/pts/model/deepar/deepar_network.py +++ b/pts/model/deepar/deepar_network.py @@ -393,8 +393,7 @@ class DeepARPredictionNetwork(DeepARNetwork): return samples.reshape(( (-1, self.num_parallel_samples) + (self.prediction_length,) - + self.target_shape - ) + + self.target_shape) ) # noinspection PyMethodOverriding,PyPep8Naming @@ -411,7 +410,6 @@ class DeepARPredictionNetwork(DeepARNetwork): Predicts samples, all tensors should have NTC layout. Parameters ---------- - F feat_static_cat : (batch_size, num_features) feat_static_real : (batch_size, num_features) past_time_feat : (batch_size, history_length, num_features) diff --git a/pts/model/estimator.py b/pts/model/estimator.py index fcc5f6d..a56647e 100644 --- a/pts/model/estimator.py +++ b/pts/model/estimator.py @@ -79,19 +79,18 @@ class PTSEstimator(Estimator): """ pass - # @abstractmethod - # def create_predictor( - # self, transformation: Transformation, trained_network: nn.Module - # ) -> Predictor: - # """ - # Create and return a predictor object. + @abstractmethod + def create_predictor( + self, transformation: Transformation, trained_network: nn.Module) -> Predictor: + """ + Create and return a predictor object. - # Returns - # ------- - # Predictor - # A predictor wrapping a `nn.Module` used for inference. - # """ - # pass + Returns + ------- + Predictor + A predictor wrapping a `nn.Module` used for inference. + """ + pass def train_model(self, training_data: Dataset) -> TrainOutput: transformation = self.create_transformation() @@ -119,7 +118,7 @@ class PTSEstimator(Estimator): return TrainOutput( transformation=transformation, trained_net=trained_net, - predictor=None#self.create_predictor(transformation, trained_net), + predictor=self.create_predictor(transformation, trained_net), ) def train(self, training_data: Dataset) -> Predictor: diff --git a/pts/model/predictor.py b/pts/model/predictor.py index b9db217..84942b6 100644 --- a/pts/model/predictor.py +++ b/pts/model/predictor.py @@ -1,6 +1,9 @@ from abc import ABC, abstractmethod from typing import Iterator +import torch +import torch.nn as nn + from pts.dataset import Dataset from .forecast import Forecast @@ -14,3 +17,6 @@ class Predictor(ABC): @abstractmethod def predict(self, dataset: Dataset, **kwargs) -> Iterator[Forecast]: pass + +class PTSPredictor(Predictor): + BlockType = nn.Module diff --git a/pts/model/utils.py b/pts/model/utils.py index f2d8322..cda860f 100644 --- a/pts/model/utils.py +++ b/pts/model/utils.py @@ -6,3 +6,7 @@ import torch.nn as nn def get_module_forward_input_names(module: nn.Module): params = inspect.signature(module.forward).parameters return list(params) + + +def copy_parameters(net_source: nn.Module, net_dest: nn.Module) -> None: + net_dest.load_state_dict(net_source.state_dict())