diff --git a/pts/dataset/__init__.py b/pts/dataset/__init__.py index 2f7e766..e9d9e6c 100644 --- a/pts/dataset/__init__.py +++ b/pts/dataset/__init__.py @@ -1 +1 @@ -from .loader import TransformedIterableDataset \ No newline at end of file +from .loader import TransformedIterableDataset diff --git a/pts/dataset/repository/__init__.py b/pts/dataset/repository/__init__.py index b74b6ee..1798039 100644 --- a/pts/dataset/repository/__init__.py +++ b/pts/dataset/repository/__init__.py @@ -1 +1 @@ -from .datasets import dataset_recipes \ No newline at end of file +from .datasets import dataset_recipes diff --git a/pts/feature/fourier_date_feature.py b/pts/feature/fourier_date_feature.py index 6912acd..18deaa7 100644 --- a/pts/feature/fourier_date_feature.py +++ b/pts/feature/fourier_date_feature.py @@ -47,6 +47,7 @@ class FourierDateFeatures(TimeFeature): steps = [x * 2.0 * np.pi / num_values for x in values] return np.vstack([np.cos(steps), np.sin(steps)]) + def fourier_time_features_from_frequency(freq_str: str) -> List[TimeFeature]: offset = to_offset(freq_str) multiple, granularity = offset.n, offset.name @@ -66,4 +67,4 @@ def fourier_time_features_from_frequency(freq_str: str) -> List[TimeFeature]: feature_classes: List[TimeFeature] = [ FourierDateFeatures(freq=freq) for freq in features[granularity] ] - return feature_classes \ No newline at end of file + return feature_classes diff --git a/pts/model/deepvar/deepvar_estimator.py b/pts/model/deepvar/deepvar_estimator.py index cf6f8c2..419cc81 100644 --- a/pts/model/deepvar/deepvar_estimator.py +++ b/pts/model/deepvar/deepvar_estimator.py @@ -32,7 +32,7 @@ from pts import Trainer 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 + lags_for_fourier_time_features_from_frequency, ) from pts.model import PyTorchEstimator from pts.modules import LowRankMultivariateNormalOutput diff --git a/pts/model/tempflow/tempflow_estimator.py b/pts/model/tempflow/tempflow_estimator.py index 1003ae9..9baaa8b 100644 --- a/pts/model/tempflow/tempflow_estimator.py +++ b/pts/model/tempflow/tempflow_estimator.py @@ -26,7 +26,7 @@ from gluonts.transform import ( from pts import Trainer from pts.feature import ( fourier_time_features_from_frequency, - lags_for_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 diff --git a/pts/model/transformer/transformer_estimator.py b/pts/model/transformer/transformer_estimator.py index b3227e4..43f4f46 100644 --- a/pts/model/transformer/transformer_estimator.py +++ b/pts/model/transformer/transformer_estimator.py @@ -10,7 +10,7 @@ from gluonts.torch.modules.distribution_output import DistributionOutput 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 ( +from gluonts.transform import ( Transformation, Chain, InstanceSplitter, @@ -205,7 +205,7 @@ class TransformerEstimator(PyTorchEstimator): def create_predictor( self, transformation: Transformation, - trained_network:TransformerTrainingNetwork, + trained_network: TransformerTrainingNetwork, device: torch.device, ) -> Predictor: diff --git a/pts/model/transformer_tempflow/transformer_tempflow_network.py b/pts/model/transformer_tempflow/transformer_tempflow_network.py index e1a09a7..69e5768 100644 --- a/pts/model/transformer_tempflow/transformer_tempflow_network.py +++ b/pts/model/transformer_tempflow/transformer_tempflow_network.py @@ -3,8 +3,7 @@ from typing import List, Optional, Tuple import torch import torch.nn as nn -from gluonts.core.component import - +from gluonts.core.component import validated from pts.modules import RealNVP, MAF, FlowOutput, MeanScaler, NOPScaler diff --git a/pts/model/utils.py b/pts/model/utils.py index 209ae1b..19fbd54 100644 --- a/pts/model/utils.py +++ b/pts/model/utils.py @@ -35,7 +35,11 @@ def weighted_average( """ if weights is not None: weighted_tensor = torch.where(weights != 0, x * weights, torch.zeros_like(x)) - sum_weights = torch.clamp(weights.sum(dim=dim) if dim else weights.sum(), min=1.0) - return (weighted_tensor.sum(dim=dim) if dim else weighted_tensor.sum())/ sum_weights + sum_weights = torch.clamp( + weights.sum(dim=dim) if dim else weights.sum(), min=1.0 + ) + return ( + weighted_tensor.sum(dim=dim) if dim else weighted_tensor.sum() + ) / sum_weights else: return x.mean(dim=dim)