forecast generator

This commit is contained in:
Dr. Kashif Rasul
2019-12-09 02:35:19 +01:00
parent 1fb23c6b30
commit 02ada03cab
2 changed files with 131 additions and 16 deletions
+9 -1
View File
@@ -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
+122 -15
View File
@@ -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]:
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"])