mirror of
https://github.com/wassname/pytorch-ts.git
synced 2026-08-04 13:13:57 +08:00
updated tempflow
This commit is contained in:
@@ -27,6 +27,7 @@ from gluonts.torch.support.util import copy_parameters
|
||||
from gluonts.torch.model.predictor import PyTorchPredictor
|
||||
from gluonts.torch.modules.distribution_output import DistributionOutput
|
||||
from gluonts.model.predictor import Predictor
|
||||
|
||||
from pts.model.utils import get_module_forward_input_names
|
||||
from pts import Trainer
|
||||
from pts.model import PyTorchEstimator
|
||||
|
||||
@@ -27,6 +27,7 @@ from gluonts.transform import (
|
||||
AddAgeFeature,
|
||||
cdf_to_gaussian_forward_transform,
|
||||
)
|
||||
|
||||
from pts import Trainer
|
||||
from pts.model.utils import get_module_forward_input_names
|
||||
from pts.feature import (
|
||||
@@ -35,6 +36,7 @@ from pts.feature import (
|
||||
)
|
||||
from pts.model import PyTorchEstimator
|
||||
from pts.modules import LowRankMultivariateNormalOutput
|
||||
|
||||
from .deepvar_network import DeepVARTrainingNetwork, DeepVARPredictionNetwork
|
||||
|
||||
|
||||
|
||||
@@ -6,7 +6,6 @@ import torch.nn as nn
|
||||
|
||||
from gluonts.dataset.field_names import FieldName
|
||||
from gluonts.torch.support.util import copy_parameters
|
||||
from pts.model.utils import get_module_forward_input_names
|
||||
from gluonts.torch.model.predictor import PyTorchPredictor
|
||||
from gluonts.model.predictor import Predictor
|
||||
from gluonts.transform import (
|
||||
@@ -17,8 +16,10 @@ from gluonts.transform import (
|
||||
AddObservedValuesIndicator,
|
||||
AsNumpyArray,
|
||||
)
|
||||
|
||||
from pts.model import PyTorchEstimator
|
||||
from pts import Trainer
|
||||
from pts.model.utils import get_module_forward_input_names
|
||||
|
||||
from .lstnet_network import LSTNetTrain, LSTNetPredict
|
||||
|
||||
|
||||
@@ -14,9 +14,11 @@ from gluonts.transform import (
|
||||
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,
|
||||
|
||||
@@ -19,6 +19,7 @@ from gluonts.transform import (
|
||||
InstanceSplitter,
|
||||
ExpectedNumInstanceSampler,
|
||||
)
|
||||
|
||||
from pts.model.utils import get_module_forward_input_names
|
||||
from pts import Trainer
|
||||
from pts.model import PyTorchEstimator
|
||||
|
||||
@@ -2,15 +2,13 @@ from typing import List, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from pts import Trainer
|
||||
from pts.dataset import FieldName
|
||||
from pts.feature import (
|
||||
TimeFeature,
|
||||
fourier_time_features_from_frequency_str,
|
||||
get_fourier_lags_for_frequency,
|
||||
)
|
||||
from pts.model import PyTorchEstimator, PyTorchPredictor, copy_parameters
|
||||
from pts.transform import (
|
||||
from gluonts.dataset.field_names import FieldName
|
||||
from gluonts.time_feature import TimeFeature
|
||||
from gluonts.torch.model.predictor import PyTorchPredictor
|
||||
from gluonts.torch.support.util import copy_parameters
|
||||
from gluonts.model.predictor import Predictor
|
||||
from gluonts.torch.model.predictor import PyTorchPredictor
|
||||
from gluonts.transform import (
|
||||
Transformation,
|
||||
Chain,
|
||||
InstanceSplitter,
|
||||
@@ -24,6 +22,15 @@ from pts.transform import (
|
||||
SetFieldIfNotPresent,
|
||||
TargetDimIndicator,
|
||||
)
|
||||
|
||||
from pts import Trainer
|
||||
from pts.feature import (
|
||||
fourier_time_features_from_frequency,
|
||||
lags_for_fourier_time_features_from_frequency
|
||||
)
|
||||
from pts.model.utils import get_module_forward_input_names
|
||||
from pts.model import PyTorchEstimator
|
||||
|
||||
from .tempflow_network import TempFlowTrainingNetwork, TempFlowPredictionNetwork
|
||||
|
||||
|
||||
@@ -83,13 +90,13 @@ class TempFlowEstimator(PyTorchEstimator):
|
||||
self.lags_seq = (
|
||||
lags_seq
|
||||
if lags_seq is not None
|
||||
else get_fourier_lags_for_frequency(freq_str=freq)
|
||||
else lags_for_fourier_time_features_from_frequency(freq_str=freq)
|
||||
)
|
||||
|
||||
self.time_features = (
|
||||
time_features
|
||||
if time_features is not None
|
||||
else fourier_time_features_from_frequency_str(self.freq)
|
||||
else fourier_time_features_from_frequency(self.freq)
|
||||
)
|
||||
|
||||
self.history_length = self.context_length + max(self.lags_seq)
|
||||
@@ -181,7 +188,7 @@ class TempFlowEstimator(PyTorchEstimator):
|
||||
transformation: Transformation,
|
||||
trained_network: TempFlowTrainingNetwork,
|
||||
device: torch.device,
|
||||
) -> PyTorchPredictor:
|
||||
) -> Predictor:
|
||||
prediction_network = TempFlowPredictionNetwork(
|
||||
input_size=self.input_size,
|
||||
target_dim=self.target_dim,
|
||||
@@ -206,13 +213,14 @@ class TempFlowEstimator(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,
|
||||
prediction_length=self.prediction_length,
|
||||
device=device,
|
||||
output_transform=None,
|
||||
)
|
||||
|
||||
@@ -4,6 +4,7 @@ import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from gluonts.core.component import validated
|
||||
|
||||
from pts.model import weighted_average
|
||||
from pts.modules import RealNVP, MAF, FlowOutput, MeanScaler, NOPScaler
|
||||
|
||||
|
||||
@@ -23,6 +23,7 @@ from gluonts.transform import (
|
||||
VstackFeatures,
|
||||
SetField,
|
||||
)
|
||||
|
||||
from pts import Trainer
|
||||
from pts.model.utils import get_module_forward_input_names
|
||||
from pts.feature import (
|
||||
@@ -31,6 +32,7 @@ from pts.feature import (
|
||||
)
|
||||
from pts.model import PyTorchEstimator
|
||||
from pts.modules import StudentTOutput
|
||||
|
||||
from .transformer_network import (
|
||||
TransformerTrainingNetwork,
|
||||
TransformerPredictionNetwork,
|
||||
|
||||
@@ -2,15 +2,12 @@ from typing import List, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from pts import Trainer
|
||||
from pts.dataset import FieldName
|
||||
from pts.feature import (
|
||||
TimeFeature,
|
||||
fourier_time_features_from_frequency_str,
|
||||
get_fourier_lags_for_frequency,
|
||||
)
|
||||
from pts.model import PyTorchEstimator, PyTorchPredictor, copy_parameters
|
||||
from pts.transform import (
|
||||
from gluonts.dataset.field_names import FieldName
|
||||
from gluonts.time_feature import TimeFeature
|
||||
from gluonts.torch.support.util import copy_parameters
|
||||
from gluonts.torch.model.predictor import PyTorchPredictor
|
||||
from gluonts.model.predictor import Predictor
|
||||
from gluonts.transform import (
|
||||
Transformation,
|
||||
Chain,
|
||||
InstanceSplitter,
|
||||
@@ -24,6 +21,15 @@ from pts.transform import (
|
||||
SetFieldIfNotPresent,
|
||||
TargetDimIndicator,
|
||||
)
|
||||
|
||||
from pts import Trainer
|
||||
from pts.model import PyTorchEstimator
|
||||
from pts.model.utils import get_module_forward_input_names
|
||||
from pts.feature import (
|
||||
fourier_time_features_from_frequency,
|
||||
lags_for_fourier_time_features_from_frequency,
|
||||
)
|
||||
|
||||
from .transformer_tempflow_network import (
|
||||
TransformerTempFlowTrainingNetwork,
|
||||
TransformerTempFlowPredictionNetwork,
|
||||
@@ -94,13 +100,13 @@ class TransformerTempFlowEstimator(PyTorchEstimator):
|
||||
self.lags_seq = (
|
||||
lags_seq
|
||||
if lags_seq is not None
|
||||
else get_fourier_lags_for_frequency(freq_str=freq)
|
||||
else lags_for_fourier_time_features_from_frequency(freq_str=freq)
|
||||
)
|
||||
|
||||
self.time_features = (
|
||||
time_features
|
||||
if time_features is not None
|
||||
else fourier_time_features_from_frequency_str(self.freq)
|
||||
else fourier_time_features_from_frequency(self.freq)
|
||||
)
|
||||
|
||||
self.history_length = self.context_length + max(self.lags_seq)
|
||||
@@ -197,7 +203,7 @@ class TransformerTempFlowEstimator(PyTorchEstimator):
|
||||
transformation: Transformation,
|
||||
trained_network: TransformerTempFlowTrainingNetwork,
|
||||
device: torch.device,
|
||||
) -> PyTorchPredictor:
|
||||
) -> Predictor:
|
||||
prediction_network = TransformerTempFlowPredictionNetwork(
|
||||
input_size=self.input_size,
|
||||
target_dim=self.target_dim,
|
||||
@@ -225,13 +231,14 @@ class TransformerTempFlowEstimator(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,
|
||||
prediction_length=self.prediction_length,
|
||||
device=device,
|
||||
output_transform=None,
|
||||
)
|
||||
|
||||
@@ -3,7 +3,8 @@ from typing import List, Optional, Tuple
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from gluonts.core.component import validated
|
||||
from gluonts.core.component import
|
||||
|
||||
from pts.modules import RealNVP, MAF, FlowOutput, MeanScaler, NOPScaler
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user