mirror of
https://github.com/wassname/pytorch-ts.git
synced 2026-07-26 13:37:40 +08:00
forecast generator
This commit is contained in:
@@ -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
@@ -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"])
|
||||
|
||||
Reference in New Issue
Block a user