Files
pytorch-ts/examples/Time-Grad-Electricity.ipynb
T
2022-12-23 14:34:54 +08:00

32 KiB

In [1]:
%matplotlib inline
import matplotlib.pyplot as plt

import numpy as np
import pandas as pd

import torch
In [2]:
from gluonts.dataset.multivariate_grouper import MultivariateGrouper
from gluonts.dataset.repository.datasets import dataset_recipes, get_dataset
from gluonts.evaluation.backtest import make_evaluation_predictions
from gluonts.evaluation import MultivariateEvaluator
In [3]:
from pts.model.tempflow import TempFlowEstimator
from pts.model.time_grad import TimeGradEstimator
from pts.model.transformer_tempflow import TransformerTempFlowEstimator
from pts import Trainer
In [4]:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
In [5]:
def plot(target, forecast, prediction_length, prediction_intervals=(50.0, 90.0), color='g', fname=None):
    label_prefix = ""
    rows = 4
    cols = 4
    fig, axs = plt.subplots(rows, cols, figsize=(24, 24))
    axx = axs.ravel()
    seq_len, target_dim = target.shape
    
    ps = [50.0] + [
            50.0 + f * c / 2.0 for c in prediction_intervals for f in [-1.0, +1.0]
        ]
        
    percentiles_sorted = sorted(set(ps))
    
    def alpha_for_percentile(p):
        return (p / 100.0) ** 0.3
        
    for dim in range(0, min(rows * cols, target_dim)):
        ax = axx[dim]

        target[-2 * prediction_length :][dim].plot(ax=ax)
        
        ps_data = [forecast.quantile(p / 100.0)[:,dim] for p in percentiles_sorted]
        i_p50 = len(percentiles_sorted) // 2
        
        p50_data = ps_data[i_p50]
        p50_series = pd.Series(data=p50_data, index=forecast.index)
        p50_series.plot(color=color, ls="-", label=f"{label_prefix}median", ax=ax)
        
        for i in range(len(percentiles_sorted) // 2):
            ptile = percentiles_sorted[i]
            alpha = alpha_for_percentile(ptile)
            ax.fill_between(
                forecast.index,
                ps_data[i],
                ps_data[-i - 1],
                facecolor=color,
                alpha=alpha,
                interpolate=True,
            )
            # Hack to create labels for the error intervals.
            # Doesn't actually plot anything, because we only pass a single data point
            pd.Series(data=p50_data[:1], index=forecast.index[:1]).plot(
                color=color,
                alpha=alpha,
                linewidth=10,
                label=f"{label_prefix}{100 - ptile * 2}%",
                ax=ax,
            )

    legend = ["observations", "median prediction"] + [f"{k}% prediction interval" for k in prediction_intervals][::-1]    
    axx[0].legend(legend, loc="upper left")
    
    if fname is not None:
        plt.savefig(fname, bbox_inches='tight', pad_inches=0.05)
In [6]:
print(f"Available datasets: {list(dataset_recipes.keys())}")
Available datasets: ['constant', 'exchange_rate', 'solar-energy', 'electricity', 'traffic', 'exchange_rate_nips', 'electricity_nips', 'traffic_nips', 'solar_nips', 'wiki-rolling_nips', 'taxi_30min', 'kaggle_web_traffic_with_missing', 'kaggle_web_traffic_without_missing', 'kaggle_web_traffic_weekly', 'm1_yearly', 'm1_quarterly', 'm1_monthly', 'nn5_daily_with_missing', 'nn5_daily_without_missing', 'nn5_weekly', 'tourism_monthly', 'tourism_quarterly', 'tourism_yearly', 'cif_2016', 'london_smart_meters_without_missing', 'wind_farms_without_missing', 'car_parts_without_missing', 'dominick', 'fred_md', 'pedestrian_counts', 'hospital', 'covid_deaths', 'kdd_cup_2018_without_missing', 'weather', 'm3_monthly', 'm3_quarterly', 'm3_yearly', 'm3_other', 'm4_hourly', 'm4_daily', 'm4_weekly', 'm4_monthly', 'm4_quarterly', 'm4_yearly', 'm5', 'uber_tlc_daily', 'uber_tlc_hourly', 'airpassengers']
In [7]:
# exchange_rate_nips, electricity_nips, traffic_nips, solar_nips, wiki-rolling_nips, ## taxi_30min is buggy still
dataset = get_dataset("electricity_nips", regenerate=False)
In [8]:
dataset.metadata
Out [8]:
MetaData(freq='H', target=None, feat_static_cat=[CategoricalFeatureInfo(name='feat_static_cat_0', cardinality='370')], feat_static_real=[], feat_dynamic_real=[], feat_dynamic_cat=[], prediction_length=24)
In [9]:
train_grouper = MultivariateGrouper(max_target_dim=min(2000, int(dataset.metadata.feat_static_cat[0].cardinality)))

test_grouper = MultivariateGrouper(num_test_dates=int(len(dataset.test)/len(dataset.train)*2), 
                                   max_target_dim=min(2000, int(dataset.metadata.feat_static_cat[0].cardinality)))
In [10]:
dataset_train = train_grouper(dataset.train)
dataset_test = test_grouper(dataset.test)
/home/wassname/miniforge3/envs/glounts/lib/python3.9/site-packages/gluonts/dataset/multivariate_grouper.py:191: VisibleDeprecationWarning: Creating an ndarray from ragged nested sequences (which is a list-or-tuple of lists-or-tuples-or ndarrays with different lengths or shapes) is deprecated. If you meant to do this, you must specify 'dtype=object' when creating the ndarray.
  return {FieldName.TARGET: np.array([funcs(data) for data in dataset])}
In [11]:
estimator = TimeGradEstimator(
    target_dim=int(dataset.metadata.feat_static_cat[0].cardinality),
    prediction_length=dataset.metadata.prediction_length,
    context_length=dataset.metadata.prediction_length,
    cell_type='GRU',
    input_size=1484,
    freq=dataset.metadata.freq,
    loss_type='l2',
    scaling=True,
    diff_steps=100,
    beta_end=0.1,
    beta_schedule="linear",
    trainer=Trainer(device=device,
                    epochs=20,
                    learning_rate=1e-3,
                    num_batches_per_epoch=100,
                    batch_size=64,)
)
In [12]:
predictor = estimator.train(dataset_train, num_workers=0)
  0%|          | 0/99 [00:00<?, ?it/s]
  0%|          | 0/99 [00:00<?, ?it/s]
  0%|          | 0/99 [00:00<?, ?it/s]
  0%|          | 0/99 [00:00<?, ?it/s]
  0%|          | 0/99 [00:00<?, ?it/s]
  0%|          | 0/99 [00:00<?, ?it/s]
  0%|          | 0/99 [00:00<?, ?it/s]
  0%|          | 0/99 [00:00<?, ?it/s]
  0%|          | 0/99 [00:00<?, ?it/s]
  0%|          | 0/99 [00:00<?, ?it/s]
  0%|          | 0/99 [00:00<?, ?it/s]
  0%|          | 0/99 [00:00<?, ?it/s]
  0%|          | 0/99 [00:00<?, ?it/s]
  0%|          | 0/99 [00:00<?, ?it/s]
  0%|          | 0/99 [00:00<?, ?it/s]
  0%|          | 0/99 [00:00<?, ?it/s]
  0%|          | 0/99 [00:00<?, ?it/s]
---------------------------------------------------------------------------
Exception                                 Traceback (most recent call last)
/media/wassname/SGIronWolf/projects5/timeseries/pytorch-ts/examples/Time-Grad-Electricity.ipynb Cell 12 in <cell line: 1>()
----> <a href='vscode-notebook-cell:/media/wassname/SGIronWolf/projects5/timeseries/pytorch-ts/examples/Time-Grad-Electricity.ipynb#X14sZmlsZQ%3D%3D?line=0'>1</a> predictor = estimator.train(dataset_train, num_workers=0)

File /media/wassname/SGIronWolf/projects5/timeseries/pytorch-ts/pts/model/estimator.py:179, in PyTorchEstimator.train(self, training_data, validation_data, num_workers, prefetch_factor, shuffle_buffer_length, cache_data, **kwargs)
    169 def train(
    170     self,
    171     training_data: Dataset,
   (...)
    177     **kwargs,
    178 ) -> PyTorchPredictor:
--> 179     return self.train_model(
    180         training_data,
    181         validation_data,
    182         num_workers=num_workers,
    183         prefetch_factor=prefetch_factor,
    184         shuffle_buffer_length=shuffle_buffer_length,
    185         cache_data=cache_data,
    186         **kwargs,
    187     ).predictor

File /media/wassname/SGIronWolf/projects5/timeseries/pytorch-ts/pts/model/estimator.py:151, in PyTorchEstimator.train_model(self, training_data, validation_data, num_workers, prefetch_factor, shuffle_buffer_length, cache_data, **kwargs)
    133     validation_iter_dataset = TransformedIterableDataset(
    134         dataset=validation_data,
    135         transform=transformation
   (...)
    139         cache_data=cache_data,
    140     )
    141     validation_data_loader = DataLoader(
    142         validation_iter_dataset,
    143         batch_size=self.trainer.batch_size,
   (...)
    148         **kwargs,
    149     )
--> 151 self.trainer(
    152     net=trained_net,
    153     train_iter=training_data_loader,
    154     validation_iter=validation_data_loader,
    155 )
    157 return TrainOutput(
    158     transformation=transformation,
    159     trained_net=trained_net,
   (...)
    162     ),
    163 )

File /media/wassname/SGIronWolf/projects5/timeseries/pytorch-ts/pts/trainer.py:63, in Trainer.__call__(self, net, train_iter, validation_iter)
     61 # training loop
     62 with tqdm(train_iter, total=total) as it:
---> 63     for batch_no, data_entry in enumerate(it, start=1):
     64         optimizer.zero_grad()
     66         inputs = [v.to(self.device) for v in data_entry.values()]

File ~/miniforge3/envs/glounts/lib/python3.9/site-packages/tqdm/notebook.py:259, in tqdm_notebook.__iter__(self)
    257 try:
    258     it = super(tqdm_notebook, self).__iter__()
--> 259     for obj in it:
    260         # return super(tqdm...) will not catch exception
    261         yield obj
    262 # NB: except ... [ as ...] breaks IPython async KeyboardInterrupt

File ~/miniforge3/envs/glounts/lib/python3.9/site-packages/tqdm/std.py:1195, in tqdm.__iter__(self)
   1192 time = self._time
   1194 try:
-> 1195     for obj in iterable:
   1196         yield obj
   1197         # Update and possibly print the progressbar.
   1198         # Note: does not call self.update(1) for speed optimisation.

File ~/miniforge3/envs/glounts/lib/python3.9/site-packages/torch/utils/data/dataloader.py:628, in _BaseDataLoaderIter.__next__(self)
    625 if self._sampler_iter is None:
    626     # TODO(https://github.com/pytorch/pytorch/issues/76750)
    627     self._reset()  # type: ignore[call-arg]
--> 628 data = self._next_data()
    629 self._num_yielded += 1
    630 if self._dataset_kind == _DatasetKind.Iterable and \
    631         self._IterableDataset_len_called is not None and \
    632         self._num_yielded > self._IterableDataset_len_called:

File ~/miniforge3/envs/glounts/lib/python3.9/site-packages/torch/utils/data/dataloader.py:671, in _SingleProcessDataLoaderIter._next_data(self)
    669 def _next_data(self):
    670     index = self._next_index()  # may raise StopIteration
--> 671     data = self._dataset_fetcher.fetch(index)  # may raise StopIteration
    672     if self._pin_memory:
    673         data = _utils.pin_memory.pin_memory(data, self._pin_memory_device)

File ~/miniforge3/envs/glounts/lib/python3.9/site-packages/torch/utils/data/_utils/fetch.py:34, in _IterableDatasetFetcher.fetch(self, possibly_batched_index)
     32 for _ in possibly_batched_index:
     33     try:
---> 34         data.append(next(self.dataset_iter))
     35     except StopIteration:
     36         self.ended = True

File ~/miniforge3/envs/glounts/lib/python3.9/site-packages/gluonts/transform/_base.py:103, in TransformedDataset.__iter__(self)
    102 def __iter__(self) -> Iterator[DataEntry]:
--> 103     yield from self.transformation(
    104         self.base_dataset, is_train=self.is_train
    105     )

File ~/miniforge3/envs/glounts/lib/python3.9/site-packages/gluonts/transform/_base.py:124, in MapTransformation.__call__(self, data_it, is_train)
    121 def __call__(
    122     self, data_it: Iterable[DataEntry], is_train: bool
    123 ) -> Iterator:
--> 124     for data_entry in data_it:
    125         try:
    126             yield self.map_transform(data_entry.copy(), is_train)

File ~/miniforge3/envs/glounts/lib/python3.9/site-packages/gluonts/transform/_base.py:124, in MapTransformation.__call__(self, data_it, is_train)
    121 def __call__(
    122     self, data_it: Iterable[DataEntry], is_train: bool
    123 ) -> Iterator:
--> 124     for data_entry in data_it:
    125         try:
    126             yield self.map_transform(data_entry.copy(), is_train)

File ~/miniforge3/envs/glounts/lib/python3.9/site-packages/gluonts/transform/_base.py:189, in FlatMapTransformation.__call__(self, data_it, is_train)
    182     yield result
    184 if (
    185     # negative values disable the check
    186     self.max_idle_transforms > 0
    187     and num_idle_transforms > self.max_idle_transforms
    188 ):
--> 189     raise Exception(
    190         "Reached maximum number of idle transformation"
    191         " calls.\nThis means the transformation looped over"
    192         f" {self.max_idle_transforms} inputs without returning any"
    193         " output.\nThis occurred in the following"
    194         f" transformation:\n{self}"
    195     )

Exception: Reached maximum number of idle transformation calls.
This means the transformation looped over 10 inputs without returning any output.
This occurred in the following transformation:
gluonts.transform.split.InstanceSplitter(dummy_value=0.0, forecast_start_field="forecast_start", future_length=24, instance_sampler=gluonts.transform.sampler.ExpectedNumInstanceSampler(axis=-1, min_past=192, min_future=24, num_instances=1.0, total_length=585345038, n=104191), is_pad_field="is_pad", lead_time=0, output_NTC=True, past_length=192, start_field="start", target_field="target", time_series_fields=["time_feat", "observed_values"])
In [ ]:
forecast_it, ts_it = make_evaluation_predictions(dataset=dataset_test,
                                                 predictor=predictor,
                                                 num_samples=100)
In [ ]:
forecasts = list(forecast_it)
targets = list(ts_it)
In [ ]:
plot(
    target=targets[0],
    forecast=forecasts[0],
    prediction_length=dataset.metadata.prediction_length,
)
plt.show()
In [ ]:
evaluator = MultivariateEvaluator(quantiles=(np.arange(20)/20.0)[1:], 
                                  target_agg_funcs={'sum': np.sum})
In [ ]:
agg_metric, item_metrics = evaluator(targets, forecasts, num_series=len(dataset_test))
In [ ]:
print("CRPS:", agg_metric["mean_wQuantileLoss"])
print("ND:", agg_metric["ND"])
print("NRMSE:", agg_metric["NRMSE"])
print("")
print("CRPS-Sum:", agg_metric["m_sum_mean_wQuantileLoss"])
print("ND-Sum:", agg_metric["m_sum_ND"])
print("NRMSE-Sum:", agg_metric["m_sum_NRMSE"])
In [ ]: