mirror of
https://github.com/wassname/pytorch-ts.git
synced 2026-07-31 12:40:41 +08:00
optimized the imports
This commit is contained in:
+4
-3
@@ -1,8 +1,9 @@
|
||||
from pkgutil import extend_path
|
||||
from pkg_resources import get_distribution, DistributionNotFound
|
||||
from .trainer import Trainer
|
||||
from .exception import assert_pts
|
||||
|
||||
from pkg_resources import get_distribution, DistributionNotFound
|
||||
|
||||
from .exception import assert_pts
|
||||
from .trainer import Trainer
|
||||
|
||||
__path__ = extend_path(__path__, __name__) # type: ignore
|
||||
|
||||
|
||||
@@ -1,10 +1,12 @@
|
||||
import inspect
|
||||
from pydantic import BaseConfig, BaseModel, create_model
|
||||
from typing import Any
|
||||
from collections import OrderedDict
|
||||
from pts.core.serde import dump_code
|
||||
import functools
|
||||
import inspect
|
||||
from collections import OrderedDict
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from pydantic import BaseConfig, BaseModel, create_model
|
||||
|
||||
from pts.core.serde import dump_code
|
||||
|
||||
|
||||
class BaseValidatedInitializerModel(BaseModel):
|
||||
|
||||
+9
-8
@@ -1,13 +1,14 @@
|
||||
from typing import Any, Optional, cast, NamedTuple
|
||||
import json
|
||||
from functools import singledispatch
|
||||
from pts.core import fqname_for
|
||||
import numpy as np
|
||||
import textwrap
|
||||
from pydoc import locate
|
||||
import math
|
||||
import itertools
|
||||
import json
|
||||
import math
|
||||
import textwrap
|
||||
from functools import singledispatch
|
||||
from pydoc import locate
|
||||
from typing import Any, Optional, cast, NamedTuple
|
||||
|
||||
import numpy as np
|
||||
|
||||
from pts.core import fqname_for
|
||||
|
||||
bad_type_msg = textwrap.dedent(
|
||||
"""
|
||||
|
||||
+21
-21
@@ -1,24 +1,3 @@
|
||||
from .common import (
|
||||
DataEntry,
|
||||
FieldName,
|
||||
Dataset,
|
||||
MetaData,
|
||||
TrainDatasets,
|
||||
DateConstants,
|
||||
)
|
||||
from .list_dataset import ListDataset
|
||||
from .file_dataset import FileDataset
|
||||
from .loader import TrainDataLoader, InferenceDataLoader
|
||||
from .process import ProcessStartField, ProcessDataEntry
|
||||
from .utils import (
|
||||
to_pandas,
|
||||
load_datasets,
|
||||
save_datasets,
|
||||
serialize_data_entry,
|
||||
frequency_add,
|
||||
forecast_start,
|
||||
)
|
||||
from .stat import DatasetStatistics, ScaleHistogram, calculate_dataset_statistics
|
||||
from .artificial import (
|
||||
ArtificialDataset,
|
||||
ConstantDataset,
|
||||
@@ -28,5 +7,26 @@ from .artificial import (
|
||||
default_synthetic,
|
||||
generate_sf2,
|
||||
)
|
||||
from .common import (
|
||||
DataEntry,
|
||||
FieldName,
|
||||
Dataset,
|
||||
MetaData,
|
||||
TrainDatasets,
|
||||
DateConstants,
|
||||
)
|
||||
from .file_dataset import FileDataset
|
||||
from .list_dataset import ListDataset
|
||||
from .loader import TrainDataLoader, InferenceDataLoader
|
||||
from .multivariate_grouper import MultivariateGrouper
|
||||
from .process import ProcessStartField, ProcessDataEntry
|
||||
from .stat import DatasetStatistics, ScaleHistogram, calculate_dataset_statistics
|
||||
from .transformed_iterable_dataset import TransformedIterableDataset
|
||||
from .utils import (
|
||||
to_pandas,
|
||||
load_datasets,
|
||||
save_datasets,
|
||||
serialize_data_entry,
|
||||
frequency_add,
|
||||
forecast_start,
|
||||
)
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import os
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
from typing import Callable, List, NamedTuple, Optional, Tuple, Union
|
||||
|
||||
@@ -17,7 +17,6 @@ from .common import (
|
||||
DataEntry,
|
||||
)
|
||||
from .list_dataset import ListDataset
|
||||
from .stat import DatasetStatistics, calculate_dataset_statistics
|
||||
from .recipe import (
|
||||
BinaryHolidays,
|
||||
BinaryMarkovChain,
|
||||
@@ -31,6 +30,7 @@ from .recipe import (
|
||||
generate,
|
||||
take_as_list,
|
||||
)
|
||||
from .stat import DatasetStatistics, calculate_dataset_statistics
|
||||
|
||||
|
||||
class DatasetInfo(NamedTuple):
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Dict, Iterable, NamedTuple, Sized, List, Optional, Iterator
|
||||
from typing import Any, Dict, Iterable, NamedTuple, List, Optional
|
||||
|
||||
import pandas as pd
|
||||
from pydantic import BaseModel
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
import functools
|
||||
from pathlib import Path
|
||||
from typing import NamedTuple
|
||||
from typing import Iterator, List
|
||||
import glob
|
||||
import random
|
||||
from pathlib import Path
|
||||
from typing import Iterator, List
|
||||
from typing import NamedTuple
|
||||
|
||||
import rapidjson as json
|
||||
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from typing import Iterable
|
||||
import random
|
||||
from typing import Iterable
|
||||
|
||||
from .common import DataEntry, Dataset, SourceContext
|
||||
from .process import ProcessDataEntry
|
||||
|
||||
@@ -3,13 +3,12 @@ from collections import defaultdict
|
||||
from typing import Any, Dict, Iterable, Iterator, List, Optional # noqa: F401
|
||||
|
||||
import numpy as np
|
||||
|
||||
# Third-party imports
|
||||
import torch
|
||||
|
||||
from pts.transform.transform import Transformation
|
||||
# First-party imports
|
||||
from .common import DataEntry, Dataset
|
||||
from pts.transform import Transformation
|
||||
|
||||
DataBatch = Dict[str, Any]
|
||||
|
||||
|
||||
@@ -13,9 +13,10 @@
|
||||
|
||||
# Standard library imports
|
||||
import logging
|
||||
from typing import Callable, Optional
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
from typing import Callable, Optional
|
||||
|
||||
# First-party imports
|
||||
from .common import DataEntry, Dataset, FieldName, DateConstants
|
||||
|
||||
@@ -18,14 +18,13 @@ large files in GluonTS master.
|
||||
"""
|
||||
import json
|
||||
import os
|
||||
import tarfile
|
||||
import shutil
|
||||
import tarfile
|
||||
from pathlib import Path
|
||||
from typing import NamedTuple, Optional
|
||||
from urllib import request
|
||||
|
||||
from pts.dataset import FileDataset, FieldName
|
||||
|
||||
from ._util import save_to_file, to_dict, metadata
|
||||
|
||||
|
||||
|
||||
@@ -23,7 +23,6 @@ from typing import List, NamedTuple, Optional
|
||||
import pandas as pd
|
||||
|
||||
from pts.dataset import frequency_add
|
||||
|
||||
from ._util import save_to_file, to_dict, metadata
|
||||
|
||||
|
||||
|
||||
@@ -11,12 +11,12 @@
|
||||
# express or implied. See the License for the specific language governing
|
||||
# permissions and limitations under the License.
|
||||
|
||||
from pathlib import Path
|
||||
import os
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import pandas as pd
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
|
||||
from ._util import save_to_file, to_dict, metadata
|
||||
|
||||
|
||||
@@ -17,11 +17,10 @@ from functools import partial
|
||||
from pathlib import Path
|
||||
|
||||
from pts.dataset import ConstantDataset, TrainDatasets, load_datasets
|
||||
|
||||
from ._artificial import generate_artificial_dataset
|
||||
from ._gp_copula_2019 import generate_gp_copula_dataset
|
||||
from ._lstnet import generate_lstnet_dataset
|
||||
from ._m4 import generate_m4_dataset
|
||||
from ._gp_copula_2019 import generate_gp_copula_dataset
|
||||
from ._util import get_download_path
|
||||
|
||||
m4_freq = "Hourly"
|
||||
|
||||
+1
-1
@@ -5,8 +5,8 @@ from typing import Any, List, NamedTuple, Optional, Set
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
|
||||
from .common import FieldName
|
||||
from pts.exception import assert_pts
|
||||
from .common import FieldName
|
||||
|
||||
|
||||
class ScaleHistogram:
|
||||
|
||||
@@ -1,12 +1,10 @@
|
||||
import itertools
|
||||
from typing import Dict, Iterable, Iterator, Optional
|
||||
import random
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from pts.transform import Transformation
|
||||
|
||||
from pts.transform.transform import Transformation
|
||||
from .common import DataEntry, Dataset
|
||||
|
||||
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from pathlib import Path
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
|
||||
@@ -1,2 +1,2 @@
|
||||
from .evaluator import Evaluator, MultivariateEvaluator
|
||||
from .backtest import make_evaluation_predictions, backtest_metrics
|
||||
from .evaluator import Evaluator, MultivariateEvaluator
|
||||
|
||||
@@ -5,16 +5,15 @@ from typing import Dict, Iterator, NamedTuple, Optional, Tuple, Union
|
||||
# Third-party imports
|
||||
import pandas as pd
|
||||
|
||||
# First-party imports
|
||||
from pts.transform import AdhocTransform, TransformedDataset
|
||||
from pts.dataset import (
|
||||
DataEntry,
|
||||
Dataset,
|
||||
InferenceDataLoader,
|
||||
DatasetStatistics,
|
||||
calculate_dataset_statistics,
|
||||
)
|
||||
from pts.model import Estimator, PTSEstimator, PTSPredictor, Predictor, Forecast
|
||||
from pts.model import Estimator, Predictor, Forecast
|
||||
# First-party imports
|
||||
from pts.transform import AdhocTransform, TransformedDataset
|
||||
from .evaluator import Evaluator
|
||||
|
||||
|
||||
|
||||
@@ -10,15 +10,14 @@ from typing import (
|
||||
Union,
|
||||
Callable,
|
||||
)
|
||||
from tqdm import tqdm
|
||||
|
||||
# Third-party imports
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
from tqdm import tqdm
|
||||
|
||||
from pts.model import Quantile, Forecast
|
||||
from pts.feature import get_seasonality
|
||||
from pts.model import Quantile, Forecast
|
||||
|
||||
|
||||
class Evaluator:
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
from .holiday import SPECIAL_DATE_FEATURES, SpecialDateFeatureSet
|
||||
from .lag import get_lags_for_frequency, get_fourier_lags_for_frequency
|
||||
from .time_feature import (
|
||||
DayOfMonth,
|
||||
@@ -12,5 +13,4 @@ from .time_feature import (
|
||||
time_features_from_frequency_str,
|
||||
fourier_time_features_from_frequency_str,
|
||||
)
|
||||
from .holiday import SPECIAL_DATE_FEATURES, SpecialDateFeatureSet
|
||||
from .utils import get_granularity, get_seasonality
|
||||
|
||||
+1
-1
@@ -12,7 +12,7 @@
|
||||
# permissions and limitations under the License.
|
||||
|
||||
# Standard library imports
|
||||
from typing import List, Optional, Tuple
|
||||
from typing import List, Optional
|
||||
|
||||
# Third-party imports
|
||||
import numpy as np
|
||||
|
||||
@@ -4,6 +4,7 @@ from typing import List
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
from pandas.tseries.frequencies import to_offset
|
||||
|
||||
from pts.core.component import validated
|
||||
from .utils import get_granularity
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import re
|
||||
from typing import Tuple
|
||||
from functools import lru_cache
|
||||
from typing import Tuple
|
||||
|
||||
|
||||
def get_granularity(freq_str: str) -> Tuple[int, str]:
|
||||
|
||||
@@ -1,16 +1,18 @@
|
||||
from typing import List, Optional
|
||||
|
||||
import numpy as np
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from pts import Trainer
|
||||
from pts.dataset import FieldName
|
||||
from pts.feature import (
|
||||
TimeFeature,
|
||||
get_lags_for_frequency,
|
||||
time_features_from_frequency_str,
|
||||
)
|
||||
from pts.model import PTSEstimator, Predictor, PTSPredictor, copy_parameters
|
||||
from pts.modules import DistributionOutput, StudentTOutput
|
||||
from pts.transform import (
|
||||
Transformation,
|
||||
Chain,
|
||||
@@ -24,10 +26,6 @@ from pts.transform import (
|
||||
InstanceSplitter,
|
||||
ExpectedNumInstanceSampler,
|
||||
)
|
||||
from pts.dataset import FieldName
|
||||
from pts.model import PTSEstimator, Predictor, PTSPredictor, copy_parameters
|
||||
from pts.modules import DistributionOutput, StudentTOutput
|
||||
|
||||
from .deepar_network import DeepARTrainingNetwork, DeepARPredictionNetwork
|
||||
|
||||
|
||||
|
||||
@@ -1,13 +1,13 @@
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.distributions import Distribution
|
||||
|
||||
import numpy as np
|
||||
from pts.core.component import validated
|
||||
from pts.modules import DistributionOutput, MeanScaler, NOPScaler, FeatureEmbedder
|
||||
from pts.model import weighted_average
|
||||
from pts.modules import DistributionOutput, MeanScaler, NOPScaler, FeatureEmbedder
|
||||
|
||||
|
||||
def prod(xs):
|
||||
|
||||
@@ -1,14 +1,16 @@
|
||||
from typing import List, Optional
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
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 PTSEstimator, PTSPredictor, copy_parameters
|
||||
from pts.modules import DistributionOutput, LowRankMultivariateNormalOutput
|
||||
from pts.dataset import FieldName
|
||||
from pts.transform import (
|
||||
Transformation,
|
||||
Chain,
|
||||
@@ -25,12 +27,6 @@ from pts.transform import (
|
||||
SetFieldIfNotPresent,
|
||||
TargetDimIndicator,
|
||||
)
|
||||
from pts.feature import (
|
||||
TimeFeature,
|
||||
fourier_time_features_from_frequency_str,
|
||||
get_fourier_lags_for_frequency,
|
||||
)
|
||||
|
||||
from .deepvar_network import DeepVARTrainingNetwork, DeepVARPredictionNetwork
|
||||
|
||||
|
||||
|
||||
@@ -2,12 +2,10 @@ from typing import List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.distributions import Distribution
|
||||
|
||||
import numpy as np
|
||||
from pts.core.component import validated
|
||||
from pts.modules import DistributionOutput, MeanScaler, NOPScaler
|
||||
from pts.model import weighted_average
|
||||
from pts.modules import DistributionOutput, MeanScaler, NOPScaler
|
||||
|
||||
|
||||
class DeepVARTrainingNetwork(nn.Module):
|
||||
|
||||
@@ -6,10 +6,9 @@ import torch
|
||||
import torch.nn as nn
|
||||
from torch.utils.data import DataLoader
|
||||
|
||||
from pts import Trainer
|
||||
from pts.dataset import Dataset, TransformedIterableDataset
|
||||
from pts.transform import Transformation
|
||||
from pts import Trainer
|
||||
|
||||
from .predictor import Predictor
|
||||
from .utils import get_module_forward_input_names
|
||||
|
||||
|
||||
@@ -2,10 +2,10 @@ from abc import ABC, abstractmethod
|
||||
from enum import Enum
|
||||
from typing import Dict, List, Optional, Set, Union, Callable
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import torch
|
||||
from pydantic import BaseModel, Field
|
||||
from torch.distributions import Distribution
|
||||
|
||||
from .quantile import Quantile
|
||||
|
||||
@@ -5,10 +5,10 @@ import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from pts.core.component import validated
|
||||
from pts.dataset import InferenceDataLoader, DataEntry, FieldName
|
||||
from pts.modules import DistributionOutput
|
||||
from .forecast import Forecast, DistributionForecast, QuantileForecast, SampleForecast
|
||||
from pts.core.component import validated
|
||||
|
||||
OutputTransform = Callable[[DataEntry, np.ndarray], np.ndarray]
|
||||
|
||||
|
||||
@@ -1,2 +1,2 @@
|
||||
from .n_beats_estimator import NBEATSEstimator
|
||||
from .n_beats_ensemble import NBEATSEnsembleEstimator
|
||||
from .n_beats_estimator import NBEATSEstimator
|
||||
|
||||
@@ -1,16 +1,15 @@
|
||||
from typing import List, Optional, Iterator
|
||||
from itertools import product
|
||||
import copy
|
||||
import logging
|
||||
from itertools import product
|
||||
from typing import List, Optional, Iterator
|
||||
|
||||
import numpy as np
|
||||
|
||||
from pts.model import Predictor, SampleForecast, Forecast, Estimator
|
||||
from pts import Trainer
|
||||
from pts.dataset import Dataset, FieldName
|
||||
|
||||
from .n_beats_network import VALID_LOSS_FUNCTIONS
|
||||
from pts.model import Predictor, SampleForecast, Forecast, Estimator
|
||||
from .n_beats_estimator import NBEATSEstimator
|
||||
from .n_beats_network import VALID_LOSS_FUNCTIONS
|
||||
|
||||
AGGREGATION_METHODS = "median", "mean", "none"
|
||||
|
||||
|
||||
@@ -1,11 +1,10 @@
|
||||
from typing import List, Optional
|
||||
|
||||
import numpy as np
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from pts import Trainer
|
||||
from pts.dataset import FieldName
|
||||
from pts.model import PTSEstimator, Predictor, PTSPredictor, copy_parameters
|
||||
from pts.transform import (
|
||||
InstanceSplitter,
|
||||
@@ -14,15 +13,13 @@ from pts.transform import (
|
||||
RemoveFields,
|
||||
ExpectedNumInstanceSampler,
|
||||
)
|
||||
from pts.dataset import FieldName
|
||||
|
||||
from .n_beats_network import (
|
||||
NBEATSPredictionNetwork,
|
||||
NBEATSTrainingNetwork,
|
||||
VALID_N_BEATS_STACK_TYPES,
|
||||
VALID_LOSS_FUNCTIONS,
|
||||
)
|
||||
|
||||
|
||||
class NBEATSEstimator(PTSEstimator):
|
||||
def __init__(
|
||||
self,
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
from typing import List, Tuple
|
||||
|
||||
import numpy as np
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
@@ -1,19 +1,20 @@
|
||||
import pts
|
||||
import json
|
||||
from abc import ABC, abstractmethod
|
||||
from pathlib import Path
|
||||
from pydoc import locate
|
||||
from typing import Iterator, Callable, Optional
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from abc import ABC, abstractmethod
|
||||
from pydoc import locate
|
||||
from typing import Iterator, Callable, Optional
|
||||
|
||||
import pts
|
||||
from pts.core.serde import dump_json, fqname_for, load_json
|
||||
from pts.dataset import Dataset, DataEntry, InferenceDataLoader
|
||||
from pts.transform import Transformation
|
||||
from pathlib import Path
|
||||
from .forecast import Forecast
|
||||
from .forecast_generator import ForecastGenerator, SampleForecastGenerator
|
||||
from .utils import get_module_forward_input_names
|
||||
from pts.core.serde import dump_json, fqname_for, load_json
|
||||
|
||||
|
||||
OutputTransform = Callable[[DataEntry, np.ndarray], np.ndarray]
|
||||
|
||||
|
||||
@@ -4,16 +4,15 @@ import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from pts import Trainer
|
||||
from pts.dataset import FieldName
|
||||
from pts.model import PTSEstimator, PTSPredictor, copy_parameters
|
||||
from pts.modules import DistributionOutput, StudentTOutput
|
||||
from pts.dataset import FieldName
|
||||
from pts.transform import (
|
||||
Transformation,
|
||||
Chain,
|
||||
InstanceSplitter,
|
||||
ExpectedNumInstanceSampler,
|
||||
)
|
||||
|
||||
from .simple_feedforward_network import (
|
||||
SimpleFeedForwardTrainingNetwork,
|
||||
SimpleFeedForwardPredictionNetwork,
|
||||
|
||||
@@ -3,8 +3,8 @@ from typing import List
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.distributions import Distribution
|
||||
from pts.core.component import validated
|
||||
|
||||
from pts.core.component import validated
|
||||
from pts.modules import MeanScaler, NOPScaler, DistributionOutput, LambdaLayer
|
||||
|
||||
|
||||
|
||||
@@ -1,21 +1,20 @@
|
||||
from typing import List, Optional
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from pts import Trainer
|
||||
from pts.model import PTSEstimator, PTSPredictor, copy_parameters
|
||||
from pts.modules import RealNVP
|
||||
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 PTSEstimator, PTSPredictor, copy_parameters
|
||||
from pts.transform import (
|
||||
Transformation,
|
||||
Chain,
|
||||
InstanceSplitter,
|
||||
ExpectedNumInstanceSampler,
|
||||
CDFtoGaussianTransform,
|
||||
cdf_to_gaussian_forward_transform,
|
||||
RenameFields,
|
||||
AsNumpyArray,
|
||||
ExpandDimArray,
|
||||
@@ -25,12 +24,6 @@ from pts.transform import (
|
||||
SetFieldIfNotPresent,
|
||||
TargetDimIndicator,
|
||||
)
|
||||
from pts.feature import (
|
||||
TimeFeature,
|
||||
fourier_time_features_from_frequency_str,
|
||||
get_fourier_lags_for_frequency,
|
||||
)
|
||||
|
||||
from .tempflow_network import TempFlowTrainingNetwork, TempFlowPredictionNetwork
|
||||
|
||||
|
||||
|
||||
@@ -3,10 +3,9 @@ from typing import List, Optional, Tuple, Union
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
import numpy as np
|
||||
from pts.core.component import validated
|
||||
from pts.modules import RealNVP, MAF, FlowOutput, MeanScaler, NOPScaler
|
||||
from pts.model import weighted_average
|
||||
from pts.modules import RealNVP, MAF, FlowOutput, MeanScaler, NOPScaler
|
||||
|
||||
|
||||
class TempFlowTrainingNetwork(nn.Module):
|
||||
|
||||
@@ -1,39 +1,31 @@
|
||||
from typing import List, Optional
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from pts import Trainer
|
||||
from pts.model import PTSEstimator, Predictor, PTSPredictor, copy_parameters
|
||||
from pts.modules import DistributionOutput, StudentTOutput
|
||||
from pts.dataset import FieldName
|
||||
from pts.transform import (
|
||||
Transformation,
|
||||
Chain,
|
||||
InstanceSplitter,
|
||||
ExpectedNumInstanceSampler,
|
||||
CDFtoGaussianTransform,
|
||||
cdf_to_gaussian_forward_transform,
|
||||
RenameFields,
|
||||
RemoveFields,
|
||||
AddAgeFeature,
|
||||
AsNumpyArray,
|
||||
ExpandDimArray,
|
||||
AddObservedValuesIndicator,
|
||||
AddTimeFeatures,
|
||||
VstackFeatures,
|
||||
SetFieldIfNotPresent,
|
||||
TargetDimIndicator,
|
||||
SetField,
|
||||
)
|
||||
from pts.feature import (
|
||||
TimeFeature,
|
||||
fourier_time_features_from_frequency_str,
|
||||
get_fourier_lags_for_frequency,
|
||||
)
|
||||
|
||||
from pts.model import PTSEstimator, Predictor, PTSPredictor, copy_parameters
|
||||
from pts.modules import DistributionOutput, StudentTOutput
|
||||
from pts.transform import (
|
||||
Transformation,
|
||||
Chain,
|
||||
InstanceSplitter,
|
||||
ExpectedNumInstanceSampler,
|
||||
RemoveFields,
|
||||
AddAgeFeature,
|
||||
AsNumpyArray,
|
||||
AddObservedValuesIndicator,
|
||||
AddTimeFeatures,
|
||||
VstackFeatures,
|
||||
SetField,
|
||||
)
|
||||
from .transformer_network import (
|
||||
TransformerTrainingNetwork,
|
||||
TransformerPredictionNetwork,
|
||||
|
||||
@@ -1,13 +1,10 @@
|
||||
from typing import List, Optional, Tuple, Union
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.distributions import Distribution
|
||||
|
||||
import numpy as np
|
||||
from pts.core.component import validated
|
||||
from pts.modules import DistributionOutput, MeanScaler, NOPScaler, FeatureEmbedder
|
||||
from pts.model import weighted_average
|
||||
|
||||
|
||||
def prod(xs):
|
||||
|
||||
@@ -1,21 +1,20 @@
|
||||
from typing import List, Optional
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from pts import Trainer
|
||||
from pts.model import PTSEstimator, PTSPredictor, copy_parameters
|
||||
from pts.modules import RealNVP
|
||||
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 PTSEstimator, PTSPredictor, copy_parameters
|
||||
from pts.transform import (
|
||||
Transformation,
|
||||
Chain,
|
||||
InstanceSplitter,
|
||||
ExpectedNumInstanceSampler,
|
||||
CDFtoGaussianTransform,
|
||||
cdf_to_gaussian_forward_transform,
|
||||
RenameFields,
|
||||
AsNumpyArray,
|
||||
ExpandDimArray,
|
||||
@@ -25,12 +24,6 @@ from pts.transform import (
|
||||
SetFieldIfNotPresent,
|
||||
TargetDimIndicator,
|
||||
)
|
||||
from pts.feature import (
|
||||
TimeFeature,
|
||||
fourier_time_features_from_frequency_str,
|
||||
get_fourier_lags_for_frequency,
|
||||
)
|
||||
|
||||
from .transformer_tempflow_network import TransformerTempFlowTrainingNetwork, TransformerTempFlowPredictionNetwork
|
||||
|
||||
|
||||
|
||||
@@ -1,12 +1,10 @@
|
||||
from typing import List, Optional, Tuple, Union
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
import numpy as np
|
||||
from pts.core.component import validated
|
||||
from pts.modules import RealNVP, MAF, FlowOutput, MeanScaler, NOPScaler
|
||||
from pts.model import weighted_average
|
||||
|
||||
|
||||
class TransformerTempFlowTrainingNetwork(nn.Module):
|
||||
|
||||
+1
-1
@@ -1,5 +1,5 @@
|
||||
from typing import Optional
|
||||
import inspect
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
@@ -11,7 +11,7 @@ from .distribution_output import (
|
||||
MultivariateNormalOutput,
|
||||
FlowOutput,
|
||||
)
|
||||
from .lambda_layer import LambdaLayer
|
||||
from .feature import FeatureEmbedder, FeatureAssembler
|
||||
from .scaler import MeanScaler, NOPScaler
|
||||
from .flows import RealNVP, MAF
|
||||
from .lambda_layer import LambdaLayer
|
||||
from .scaler import MeanScaler, NOPScaler
|
||||
|
||||
@@ -5,10 +5,8 @@ import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from pts.core.component import validated
|
||||
from torch.distributions import (
|
||||
Distribution,
|
||||
Normal,
|
||||
Beta,
|
||||
NegativeBinomial,
|
||||
StudentT,
|
||||
@@ -20,8 +18,8 @@ from torch.distributions import (
|
||||
AffineTransform,
|
||||
)
|
||||
|
||||
from pts.core.component import validated
|
||||
from .lambda_layer import LambdaLayer
|
||||
from .flows import RealNVP
|
||||
|
||||
|
||||
class ArgProj(nn.Module):
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from typing import Callable, List, Optional
|
||||
from typing import List, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
import copy
|
||||
import math
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from typing import Tuple
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
+2
-2
@@ -1,12 +1,12 @@
|
||||
import time
|
||||
from typing import Any, List, NamedTuple, Optional, Union
|
||||
from typing import List, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.utils.data import DataLoader
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
from tqdm import tqdm
|
||||
|
||||
from torch.utils.data import DataLoader
|
||||
|
||||
class Trainer:
|
||||
def __init__(
|
||||
|
||||
+10
-17
@@ -1,14 +1,3 @@
|
||||
from .transform import (
|
||||
Transformation,
|
||||
Chain,
|
||||
Identity,
|
||||
MapTransformation,
|
||||
SimpleTransformation,
|
||||
AdhocTransform,
|
||||
FlatMapTransformation,
|
||||
FilterTransformation,
|
||||
)
|
||||
|
||||
from .convert import (
|
||||
AsNumpyArray,
|
||||
ExpandDimArray,
|
||||
@@ -21,9 +10,7 @@ from .convert import (
|
||||
CDFtoGaussianTransform,
|
||||
cdf_to_gaussian_forward_transform,
|
||||
)
|
||||
|
||||
from .dataset import TransformedDataset
|
||||
|
||||
from .feature import (
|
||||
target_transformation_length,
|
||||
AddObservedValuesIndicator,
|
||||
@@ -31,7 +18,6 @@ from .feature import (
|
||||
AddTimeFeatures,
|
||||
AddAgeFeature,
|
||||
)
|
||||
|
||||
from .field import (
|
||||
RemoveFields,
|
||||
RenameFields,
|
||||
@@ -39,8 +25,6 @@ from .field import (
|
||||
SetFieldIfNotPresent,
|
||||
SelectFields,
|
||||
)
|
||||
|
||||
|
||||
from .sampler import (
|
||||
InstanceSampler,
|
||||
UniformSplitSampler,
|
||||
@@ -50,10 +34,19 @@ from .sampler import (
|
||||
ContinuousTimePointSampler,
|
||||
ContinuousTimeUniformSampler,
|
||||
)
|
||||
|
||||
from .split import (
|
||||
shift_timestamp,
|
||||
InstanceSplitter,
|
||||
CanonicalInstanceSplitter,
|
||||
ContinuousTimeInstanceSplitter,
|
||||
)
|
||||
from .transform import (
|
||||
Transformation,
|
||||
Chain,
|
||||
Identity,
|
||||
MapTransformation,
|
||||
SimpleTransformation,
|
||||
AdhocTransform,
|
||||
FlatMapTransformation,
|
||||
FilterTransformation,
|
||||
)
|
||||
|
||||
@@ -14,14 +14,12 @@
|
||||
from typing import Iterator, List, Tuple, Optional
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from scipy.special import erf, erfinv
|
||||
|
||||
import torch
|
||||
|
||||
from pts.exception import assert_pts
|
||||
from pts.core.component import validated
|
||||
from pts.dataset import DataEntry
|
||||
|
||||
from pts.exception import assert_pts
|
||||
from .transform import (
|
||||
SimpleTransformation,
|
||||
MapTransformation,
|
||||
|
||||
@@ -16,11 +16,11 @@ from typing import List
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
|
||||
from pts.core.component import validated
|
||||
from pts.dataset import DataEntry
|
||||
from pts.feature import TimeFeature
|
||||
from pts.core.component import validated
|
||||
from .transform import SimpleTransformation, MapTransformation
|
||||
from .split import shift_timestamp
|
||||
from .transform import SimpleTransformation, MapTransformation
|
||||
|
||||
|
||||
def target_transformation_length(
|
||||
|
||||
@@ -14,8 +14,8 @@
|
||||
from collections import Counter
|
||||
from typing import Any, Dict, List
|
||||
|
||||
from pts.dataset import DataEntry
|
||||
from pts.core.component import validated
|
||||
from pts.dataset import DataEntry
|
||||
from .transform import SimpleTransformation, MapTransformation
|
||||
|
||||
|
||||
|
||||
@@ -15,8 +15,8 @@ from abc import ABC, abstractmethod
|
||||
|
||||
import numpy as np
|
||||
|
||||
from pts.dataset.stat import ScaleHistogram
|
||||
from pts.core.component import validated
|
||||
from pts.dataset.stat import ScaleHistogram
|
||||
|
||||
|
||||
class InstanceSampler(ABC):
|
||||
|
||||
@@ -17,11 +17,10 @@ from typing import Iterator, List, Optional
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
|
||||
from pts.dataset import DataEntry, FieldName
|
||||
|
||||
from .transform import FlatMapTransformation
|
||||
from .sampler import InstanceSampler, ContinuousTimePointSampler
|
||||
from pts.core.component import validated
|
||||
from pts.dataset import DataEntry, FieldName
|
||||
from .sampler import InstanceSampler, ContinuousTimePointSampler
|
||||
from .transform import FlatMapTransformation
|
||||
|
||||
|
||||
def shift_timestamp(ts: pd.Timestamp, offset: int) -> pd.Timestamp:
|
||||
|
||||
@@ -1,11 +1,10 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Callable, Iterator, Iterable, List
|
||||
from functools import reduce
|
||||
from typing import Callable, Iterator, Iterable, List
|
||||
|
||||
from pts.core.component import validated
|
||||
|
||||
from pts.dataset import DataEntry
|
||||
|
||||
|
||||
MAX_IDLE_TRANSFORMS = 100
|
||||
|
||||
|
||||
|
||||
@@ -21,7 +21,6 @@ setup(
|
||||
'pandas',
|
||||
'scipy',
|
||||
'tqdm',
|
||||
'ujson',
|
||||
'pydantic',
|
||||
'matplotlib',
|
||||
'python-rapidjson',
|
||||
|
||||
@@ -11,9 +11,9 @@
|
||||
# express or implied. See the License for the specific language governing
|
||||
# permissions and limitations under the License.
|
||||
|
||||
import numpy as np
|
||||
# Standard library imports
|
||||
import pytest
|
||||
import numpy as np
|
||||
|
||||
# First-party imports
|
||||
from pts.dataset import ListDataset, MultivariateGrouper
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import pytest
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from pts.dataset import ProcessStartField
|
||||
|
||||
|
||||
@@ -15,12 +15,11 @@ from itertools import islice
|
||||
|
||||
import torch
|
||||
|
||||
from pts.modules import StudentTOutput
|
||||
from pts.dataset import constant_dataset, TrainDataLoader
|
||||
from pts.model.deepar import DeepAREstimator
|
||||
from pts.model import get_module_forward_input_names
|
||||
from pts import Trainer
|
||||
|
||||
from pts.dataset import constant_dataset, TrainDataLoader
|
||||
from pts.model import get_module_forward_input_names
|
||||
from pts.model.deepar import DeepAREstimator
|
||||
from pts.modules import StudentTOutput
|
||||
|
||||
ds_info, train_ds, test_ds = constant_dataset()
|
||||
freq = ds_info.metadata.freq
|
||||
|
||||
@@ -14,17 +14,17 @@
|
||||
# First-party imports
|
||||
import pytest
|
||||
|
||||
from pts import Trainer
|
||||
from pts.dataset import TrainDatasets, MultivariateGrouper
|
||||
from pts.dataset.artificial import constant_dataset
|
||||
from pts.evaluation import MultivariateEvaluator
|
||||
from pts.evaluation import backtest_metrics
|
||||
from pts.model.deepvar import DeepVAREstimator
|
||||
from pts.modules import (
|
||||
IndependentNormalOutput,
|
||||
LowRankMultivariateNormalOutput,
|
||||
MultivariateNormalOutput,
|
||||
)
|
||||
from pts.evaluation import backtest_metrics
|
||||
from pts.model.deepvar import DeepVAREstimator
|
||||
from pts.dataset import TrainDatasets, MultivariateGrouper
|
||||
from pts import Trainer
|
||||
from pts.evaluation import MultivariateEvaluator
|
||||
|
||||
|
||||
def load_multivariate_constant_dataset():
|
||||
|
||||
@@ -16,6 +16,7 @@ import numpy as np
|
||||
import pandas as pd
|
||||
import pytest
|
||||
import torch
|
||||
from torch.distributions import Uniform
|
||||
|
||||
# First-party imports
|
||||
from pts.model import (
|
||||
@@ -24,8 +25,6 @@ from pts.model import (
|
||||
DistributionForecast,
|
||||
)
|
||||
|
||||
from torch.distributions import Uniform
|
||||
|
||||
QUANTILES = np.arange(1, 100) / 100
|
||||
SAMPLES = np.arange(101).reshape(101, 1) / 100
|
||||
START_DATE = pd.Timestamp(2017, 1, 1, 12)
|
||||
|
||||
@@ -1,13 +1,9 @@
|
||||
import pytest
|
||||
from typing import Iterable, List, Tuple
|
||||
from typing import List, Tuple
|
||||
|
||||
import numpy as np
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.nn.utils import clip_grad_norm_
|
||||
from torch.utils.data import TensorDataset, DataLoader
|
||||
from torch.optim import SGD
|
||||
from torch.distributions import (
|
||||
StudentT,
|
||||
Beta,
|
||||
@@ -17,6 +13,9 @@ from torch.distributions import (
|
||||
Independent,
|
||||
Normal,
|
||||
)
|
||||
from torch.nn.utils import clip_grad_norm_
|
||||
from torch.optim import SGD
|
||||
from torch.utils.data import TensorDataset, DataLoader
|
||||
|
||||
from pts.modules import (
|
||||
DistributionOutput,
|
||||
|
||||
@@ -1,10 +1,8 @@
|
||||
import pytest
|
||||
from itertools import chain, combinations
|
||||
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.distributions import Uniform
|
||||
|
||||
from pts.modules import FeatureEmbedder, FeatureAssembler
|
||||
|
||||
|
||||
@@ -11,14 +11,12 @@
|
||||
# express or implied. See the License for the specific language governing
|
||||
# permissions and limitations under the License.
|
||||
|
||||
import pytest
|
||||
import numpy as np
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from pts.modules import MeanScaler, NOPScaler
|
||||
|
||||
|
||||
test_cases = [
|
||||
(
|
||||
MeanScaler(),
|
||||
|
||||
@@ -17,9 +17,10 @@ from typing import Tuple
|
||||
# Third-party imports
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import torch
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from pts import transform
|
||||
# First-party imports
|
||||
from pts.dataset import (
|
||||
ProcessStartField,
|
||||
@@ -29,7 +30,6 @@ from pts.dataset import (
|
||||
calculate_dataset_statistics,
|
||||
ScaleHistogram,
|
||||
)
|
||||
from pts import transform
|
||||
from pts.feature import time_feature
|
||||
|
||||
FREQ = "1D"
|
||||
|
||||
Reference in New Issue
Block a user