mirror of
https://github.com/wassname/pytorch-ts.git
synced 2026-07-25 13:30:12 +08:00
updated nbeats
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user