updated nbeats

This commit is contained in:
Dr. Kashif Rasul
2021-01-01 22:33:40 +01:00
parent c7e603be3e
commit 4fc70dc7b5
3 changed files with 18 additions and 8 deletions
+7 -3
View File
@@ -5,9 +5,13 @@ from typing import List, Optional, Iterator
import numpy as np
from gluonts.dataset.field_names import FieldName
from gluonts.dataset.common import Dataset
from gluonts.model.predictor import Predictor
from gluonts.model.forecast import Forecast, SampleForecast
from pts.model import PyTorchEstimator
from pts import Trainer
from pts.dataset import Dataset, FieldName
from pts.model import Predictor, SampleForecast, Forecast, Estimator
from .n_beats_estimator import NBEATSEstimator
from .n_beats_network import VALID_LOSS_FUNCTIONS
@@ -89,7 +93,7 @@ class NBEATSEnsemblePredictor(Predictor):
)
class NBEATSEnsembleEstimator(Estimator):
class NBEATSEnsembleEstimator(PyTorchEstimator):
"""
An ensemble N-BEATS Estimator (approximately) as described
in the paper: https://arxiv.org/abs/1905.10437.
+10 -4
View File
@@ -3,16 +3,20 @@ from typing import List, Optional
import torch
import torch.nn as nn
from pts import Trainer
from pts.dataset import FieldName
from pts.model import PyTorchEstimator, Predictor, PyTorchPredictor, copy_parameters
from pts.transform import (
from gluonts.dataset.field_names import FieldName
from gluonts.model.predictor import Predictor
from gluonts.torch.model.predictor import PyTorchPredictor
from gluonts.torch.support.util import copy_parameters
from gluonts.transform import (
InstanceSplitter,
Transformation,
Chain,
RemoveFields,
ExpectedNumInstanceSampler,
)
from pts import Trainer
from pts.model import PyTorchEstimator
from pts.model.utils import get_module_forward_input_names
from .n_beats_network import (
NBEATSPredictionNetwork,
NBEATSTrainingNetwork,
@@ -179,9 +183,11 @@ class NBEATSEstimator(PyTorchEstimator):
).to(device)
copy_parameters(trained_network, prediction_network)
input_names = get_module_forward_input_names(prediction_network)
return PyTorchPredictor(
input_transform=transformation,
input_names=input_names,
prediction_net=prediction_network,
batch_size=self.trainer.batch_size,
freq=self.freq,
+1 -1
View File
@@ -5,7 +5,7 @@ import torch
import torch.nn as nn
import torch.nn.functional as F
from pts.feature import get_seasonality
from gluonts.time_feature import get_seasonality
VALID_N_BEATS_STACK_TYPES = "G", "S", "T"
VALID_LOSS_FUNCTIONS = "sMAPE", "MASE", "MAPE"