From 7baca9339f5a2a632e3d547895f57619fbc88fef Mon Sep 17 00:00:00 2001 From: "Dr. Kashif Rasul" Date: Fri, 6 Dec 2019 16:38:16 +0100 Subject: [PATCH] forecast and predictor --- pts/dataset/__init__.py | 2 +- pts/model/__init__.py | 2 +- pts/model/forecast.py | 143 ++++++++++++++++---------------- pts/model/forecast_generator.py | 43 ++++++++++ pts/model/predictor.py | 2 + 5 files changed, 118 insertions(+), 74 deletions(-) create mode 100644 pts/model/forecast_generator.py diff --git a/pts/dataset/__init__.py b/pts/dataset/__init__.py index 3675109..82a9193 100644 --- a/pts/dataset/__init__.py +++ b/pts/dataset/__init__.py @@ -1,6 +1,6 @@ from .common import DataEntry, FieldName, Dataset from .list_dataset import ListDataset -from .loader import TrainDataLoader +from .loader import TrainDataLoader, InferenceDataLoader from .sampler import ( InstanceSampler, BucketInstanceSampler, diff --git a/pts/model/__init__.py b/pts/model/__init__.py index 3338ed7..6c87161 100644 --- a/pts/model/__init__.py +++ b/pts/model/__init__.py @@ -1,5 +1,5 @@ from .estimator import Estimator, PTSEstimator -from .forecast import Forecast +from .forecast import Forecast, SampleForecast, QuantileForecast from .predictor import Predictor from .quantile import Quantile from .utils import get_module_forward_input_names, copy_parameters \ No newline at end of file diff --git a/pts/model/forecast.py b/pts/model/forecast.py index 87a3b72..d98f179 100644 --- a/pts/model/forecast.py +++ b/pts/model/forecast.py @@ -322,84 +322,83 @@ class QuantileForecast(Forecast): ) -# class DistributionForecast(Forecast): -# """ -# A `Forecast` object that uses a distribution directly. -# This can for instance be used to represent marginal probability -# distributions for each time point -- although joint distributions are -# also possible, e.g. when using MultiVariateGaussian). +class DistributionForecast(Forecast): + """ + A `Forecast` object that uses a distribution directly. + This can for instance be used to represent marginal probability + distributions for each time point -- although joint distributions are + also possible, e.g. when using MultiVariateGaussian). -# Parameters -# ---------- -# distribution -# Distribution object. This should represent the entire prediction -# length, i.e., if we draw `num_samples` samples from the distribution, -# the sample shape should be + Parameters + ---------- + distribution + Distribution object. This should represent the entire prediction + length, i.e., if we draw `num_samples` samples from the distribution, + the sample shape should be -# samples = trans_dist.sample(num_samples) -# samples.shape -> (num_samples, prediction_length) + samples = trans_dist.sample(num_samples) + samples.shape -> (num_samples, prediction_length) -# start_date -# start of the forecast -# freq -# forecast frequency -# info -# additional information that the forecaster may provide e.g. estimated -# parameters, number of iterations ran etc. -# """ -# @validated() -# def __init__( -# self, -# distribution: Distribution, -# start_date, -# freq, -# item_id: Optional[str] = None, -# info: Optional[Dict] = None, -# ): -# self.distribution = distribution -# self.shape = (self.distribution.batch_shape + -# self.distribution.event_shape) -# self.prediction_length = self.shape[0] -# self.item_id = item_id -# self.info = info + start_date + start of the forecast + freq + forecast frequency + info + additional information that the forecaster may provide e.g. estimated + parameters, number of iterations ran etc. + """ + def __init__( + self, + distribution: Distribution, + start_date, + freq, + item_id: Optional[str] = None, + info: Optional[Dict] = None, + ): + self.distribution = distribution + self.shape = (self.distribution.batch_shape + + self.distribution.event_shape) + self.prediction_length = self.shape[0] + self.item_id = item_id + self.info = info -# assert isinstance( -# start_date, -# pd.Timestamp), "start_date should be a pandas Timestamp object" -# self.start_date = start_date + assert isinstance( + start_date, + pd.Timestamp), "start_date should be a pandas Timestamp object" + self.start_date = start_date -# assert isinstance(freq, str), "freq should be a string" -# self.freq = freq -# self._mean = None + assert isinstance(freq, str), "freq should be a string" + self.freq = freq + self._mean = None -# @property -# def mean(self): -# """ -# Forecast mean. -# """ -# if self._mean is not None: -# return self._mean -# else: -# self._mean = self.distribution.mean.asnumpy() -# return self._mean + @property + def mean(self): + """ + Forecast mean. + """ + if self._mean is not None: + return self._mean + else: + self._mean = self.distribution.mean.asnumpy() + return self._mean -# @property -# def mean_ts(self): -# """ -# Forecast mean, as a pandas.Series object. -# """ -# return pd.Series(self.index, self.mean) + @property + def mean_ts(self): + """ + Forecast mean, as a pandas.Series object. + """ + return pd.Series(self.index, self.mean) -# def quantile(self, level): -# level = Quantile.parse(level).value -# q = self.distribution.quantile(mx.nd.array([level])).asnumpy()[0] -# return q + def quantile(self, level): + level = Quantile.parse(level).value + q = self.distribution.quantile(mx.nd.array([level])).asnumpy()[0] + return q -# def to_sample_forecast(self, num_samples: int = 200) -> SampleForecast: -# return SampleForecast( -# samples=self.distribution.sample(num_samples), -# start_date=self.start_date, -# freq=self.freq, -# item_id=self.item_id, -# info=self.info, -# ) + def to_sample_forecast(self, num_samples: int = 200) -> SampleForecast: + return SampleForecast( + samples=self.distribution.sample(num_samples), + start_date=self.start_date, + freq=self.freq, + item_id=self.item_id, + info=self.info, + ) diff --git a/pts/model/forecast_generator.py b/pts/model/forecast_generator.py new file mode 100644 index 0000000..275dab6 --- /dev/null +++ b/pts/model/forecast_generator.py @@ -0,0 +1,43 @@ +from abc import ABC, abstractmethod +from typing import Any, Callable, Iterator, List, Optional + +import numpy as np +import torch +import torch.nn as nn + +from pts.dataset import InferenceDataLoader, DataEntry +from pts.model import Forecast, DistributionForecast +from pts.modules import DistributionOutput + +OutputTransform = Callable[[DataEntry, np.ndarray], np.ndarray] + + +class ForecastGenerator(ABC): + """ + Classes used to bring the output of a network into a class. + """ + @abstractmethod + def __call__(self, + inference_data_loader: InferenceDataLoader, + prediction_net: nn.Module, + input_names: List[str], + freq: str, + output_transform: Optional[OutputTransform], + num_samples: Optional[int], + **kwargs) -> Iterator[Forecast]: + pass + + +class DistributionForecastGenerator(ForecastGenerator): + def __init__(self, distr_output: DistributionOutput) -> None: + self.distr_output = distr_output + + def __call__(self, + inference_data_loader: InferenceDataLoader, + prediction_net: nn.Module, + input_names: List[str], + freq: str, + output_transform: Optional[OutputTransform], + num_samples: Optional[int], + **kwargs) -> Iterator[DistributionForecast]: + \ No newline at end of file diff --git a/pts/model/predictor.py b/pts/model/predictor.py index 84942b6..14ec5cf 100644 --- a/pts/model/predictor.py +++ b/pts/model/predictor.py @@ -20,3 +20,5 @@ class Predictor(ABC): class PTSPredictor(Predictor): BlockType = nn.Module + +