From 02ada03cab7cde1c59982dc40d9b9a6987e86fe8 Mon Sep 17 00:00:00 2001 From: "Dr. Kashif Rasul" Date: Mon, 9 Dec 2019 02:35:19 +0100 Subject: [PATCH] forecast generator --- pts/model/forecast.py | 10 ++- pts/model/forecast_generator.py | 137 ++++++++++++++++++++++++++++---- 2 files changed, 131 insertions(+), 16 deletions(-) diff --git a/pts/model/forecast.py b/pts/model/forecast.py index 09b9966..fc1dc7b 100644 --- a/pts/model/forecast.py +++ b/pts/model/forecast.py @@ -2,6 +2,7 @@ from abc import ABC, abstractmethod from enum import Enum from typing import Dict, List, Optional, Set, Union, Callable +from pydantic import BaseModel, Field import numpy as np import pandas as pd import torch @@ -16,10 +17,17 @@ class OutputType(str, Enum): quantiles = "quantiles" -class Config: +class Config(BaseModel): + num_samples: int = Field(100, alias="num_eval_samples") output_types: Set[OutputType] = {"quantiles", "mean"} + # FIXME: validate list elements quantiles: List[str] = ["0.1", "0.5", "0.9"] + class Config: + allow_population_by_field_name = True + # store additional fields + extra = "allow" + class Forecast(ABC): start_date: pd.Timestamp diff --git a/pts/model/forecast_generator.py b/pts/model/forecast_generator.py index 275dab6..4826736 100644 --- a/pts/model/forecast_generator.py +++ b/pts/model/forecast_generator.py @@ -6,25 +6,49 @@ import torch import torch.nn as nn from pts.dataset import InferenceDataLoader, DataEntry -from pts.model import Forecast, DistributionForecast +from pts.model import Forecast, DistributionForecast, QuantileForecast, SampleForecast from pts.modules import DistributionOutput OutputTransform = Callable[[DataEntry, np.ndarray], np.ndarray] +def _extract_instances(x: Any) -> Any: + """ + Helper function to extract individual instances from batched + mxnet results. + + For a tensor `a` + _extract_instances(a) -> [a[0], a[1], ...] + + For (nested) tuples of tensors `(a, (b, c))` + _extract_instances((a, (b, c)) -> [(a[0], (b[0], c[0])), (a[1], (b[1], c[1])), ...] + """ + if isinstance(x, (np.ndarray, torch.Tensor)): + for i in range(x.shape[0]): + # yield x[i: i + 1] + yield x[i] + elif isinstance(x, tuple): + for m in zip(*[_extract_instances(y) for y in x]): + yield tuple([r for r in m]) + elif isinstance(x, list): + for m in zip(*[_extract_instances(y) for y in x]): + yield [r for r in m] + elif x is None: + while True: + yield None + else: + assert False + + 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, + 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]: + num_samples: Optional[int], **kwargs) -> Iterator[Forecast]: pass @@ -32,12 +56,95 @@ 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, + 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], + num_samples: Optional[int], **kwargs) -> Iterator[DistributionForecast]: - \ No newline at end of file + for batch in inference_data_loader: + inputs = [batch[k] for k in input_names] + outputs = prediction_net(*inputs) + if output_transform is not None: + outputs = output_transform(batch, outputs) + + distributions = [ + self.distr_output.distribution(*u) + for u in _extract_instances(outputs) + ] + + i = -1 + for i, distr in enumerate(distributions): + yield DistributionForecast( + distr, + start_date=batch["forecast_start"][i], + freq=freq, + item_id=batch[FieldName.ITEM_ID][i] + if FieldName.ITEM_ID in batch else None, + info=batch["info"][i] if "info" in batch else None, + ) + assert i + 1 == len(batch["forecast_start"]) + + +class QuantileForecastGenerator(ForecastGenerator): + def __init__(self, quantiles: List[str]) -> None: + self.quantiles = quantiles + + def __call__(self, inference_data_loader: InferenceDataLoader, + prediction_net: BlockType, input_names: List[str], freq: str, + output_transform: Optional[OutputTransform], + num_samples: Optional[int], **kwargs) -> Iterator[Forecast]: + for batch in inference_data_loader: + inputs = [batch[k] for k in input_names] + outputs = prediction_net(*inputs).numpy() + if output_transform is not None: + outputs = output_transform(batch, outputs) + + i = -1 + for i, output in enumerate(outputs): + yield QuantileForecast( + output, + start_date=batch["forecast_start"][i], + freq=freq, + item_id=batch[FieldName.ITEM_ID][i] + if FieldName.ITEM_ID in batch else None, + info=batch["info"][i] if "info" in batch else None, + forecast_keys=self.quantiles, + ) + assert i + 1 == len(batch["forecast_start"]) + + +class SampleForecastGenerator(ForecastGenerator): + def __call__(self, inference_data_loader: InferenceDataLoader, + prediction_net: BlockType, input_names: List[str], freq: str, + output_transform: Optional[OutputTransform], + num_samples: Optional[int], **kwargs) -> Iterator[Forecast]: + for batch in inference_data_loader: + inputs = [batch[k] for k in input_names] + outputs = prediction_net(*inputs).numpy() + if output_transform is not None: + outputs = output_transform(batch, outputs) + if num_samples: + num_collected_samples = outputs[0].shape[0] + collected_samples = [outputs] + while num_collected_samples < num_samples: + outputs = prediction_net(*inputs).numpy() + if output_transform is not None: + outputs = output_transform(batch, outputs) + collected_samples.append(outputs) + num_collected_samples += outputs[0].shape[0] + outputs = [ + np.concatenate(s)[:num_samples] + for s in zip(*collected_samples) + ] + assert len(outputs[0]) == num_samples + i = -1 + for i, output in enumerate(outputs): + yield SampleForecast( + output, + start_date=batch["forecast_start"][i], + freq=freq, + item_id=batch[FieldName.ITEM_ID][i] + if FieldName.ITEM_ID in batch else None, + info=batch["info"][i] if "info" in batch else None, + ) + assert i + 1 == len(batch["forecast_start"])