mirror of
https://github.com/wassname/pytorch-ts.git
synced 2026-07-24 13:20:07 +08:00
ran isort
isort --recursive --atomic --apply pts
This commit is contained in:
+1
-1
@@ -1,2 +1,2 @@
|
||||
from .exception import assert_pts
|
||||
from .trainer import Trainer
|
||||
from .trainer import Trainer
|
||||
|
||||
@@ -1,11 +1,7 @@
|
||||
from pts.dataset.common import DataEntry, FieldName
|
||||
from pts.dataset.list_dataset import ListDataset
|
||||
from pts.dataset.sampler import (
|
||||
UniformSplitSampler,
|
||||
TestSplitSampler,
|
||||
ExpectedNumInstanceSampler,
|
||||
BucketInstanceSampler,
|
||||
)
|
||||
from pts.dataset.sampler import InstanceSampler, UniformSplitSampler, TestSplitSampler, ExpectedNumInstanceSampler, BucketInstanceSampler
|
||||
from pts.dataset.loader import DataLoader, TrainDataLoader, InferenceDataLoader
|
||||
from pts.dataset.utils import to_pandas
|
||||
from pts.dataset.loader import DataLoader, InferenceDataLoader, TrainDataLoader
|
||||
from pts.dataset.sampler import (BucketInstanceSampler,
|
||||
ExpectedNumInstanceSampler, InstanceSampler,
|
||||
TestSplitSampler, UniformSplitSampler)
|
||||
from pts.dataset.utils import to_pandas
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
from typing import Any, Dict, Sized, Iterable, NamedTuple
|
||||
from typing import Any, Dict, Iterable, NamedTuple, Sized
|
||||
|
||||
DataEntry = Dict[str, Any]
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from typing import Iterable
|
||||
|
||||
from .common import Dataset, DataEntry, SourceContext
|
||||
from .common import DataEntry, Dataset, SourceContext
|
||||
from .process import ProcessDataEntry
|
||||
|
||||
|
||||
|
||||
@@ -2,13 +2,14 @@ import itertools
|
||||
from collections import defaultdict
|
||||
from typing import Any, Dict, Iterable, Iterator, List, Optional # noqa: F401
|
||||
|
||||
import numpy as np
|
||||
# Third-party imports
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
from pts.feature.transform import Transformation
|
||||
|
||||
# First-party imports
|
||||
from .common import DataEntry, Dataset
|
||||
from pts.feature.transform import Transformation
|
||||
|
||||
DataBatch = Dict[str, Any]
|
||||
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Union, List
|
||||
from typing import List, Union
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
+1
-2
@@ -1,7 +1,6 @@
|
||||
import math
|
||||
from collections import defaultdict
|
||||
from typing import Optional
|
||||
import math
|
||||
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
from typing import List, Iterator
|
||||
from typing import Iterator, List
|
||||
|
||||
from .common import Dataset, DataEntry
|
||||
from pts.feature import Transformation, Chain
|
||||
from pts.feature import Chain, Transformation
|
||||
|
||||
from .common import DataEntry, Dataset
|
||||
|
||||
|
||||
class TransformedDataset(Dataset):
|
||||
|
||||
@@ -23,4 +23,4 @@ def to_pandas(instance: dict, freq: str = None) -> pd.Series:
|
||||
if not freq:
|
||||
freq = start.freqstr
|
||||
index = pd.date_range(start=start, periods=len(target), freq=freq)
|
||||
return pd.Series(target, index=index)
|
||||
return pd.Series(target, index=index)
|
||||
|
||||
+15
-39
@@ -1,40 +1,16 @@
|
||||
from pts.feature.time_feature import (
|
||||
TimeFeature,
|
||||
MinuteOfHour,
|
||||
HourOfDay,
|
||||
DayOfWeek,
|
||||
DayOfMonth,
|
||||
DayOfYear,
|
||||
MonthOfYear,
|
||||
WeekOfYear,
|
||||
time_features_from_frequency_str,
|
||||
)
|
||||
|
||||
from pts.feature.transform import (
|
||||
Transformation,
|
||||
Chain,
|
||||
IdentityTransformation,
|
||||
MapTransformation,
|
||||
SimpleTransformation,
|
||||
AdhocTransform,
|
||||
FlatMapTransformation,
|
||||
FilterTransformation,
|
||||
RemoveFields,
|
||||
SetField,
|
||||
AsNumpyArray,
|
||||
ExpandDimArray,
|
||||
VstackFeatures,
|
||||
ConcatFeatures,
|
||||
SwapAxes,
|
||||
ListFeatures,
|
||||
AddObservedValuesIndicator,
|
||||
RenameFields,
|
||||
AddConstFeature,
|
||||
AddTimeFeatures,
|
||||
AddAgeFeature,
|
||||
InstanceSplitter,
|
||||
CanonicalInstanceSplitter,
|
||||
SelectFields,
|
||||
)
|
||||
|
||||
from pts.feature.lag import get_lags_for_frequency
|
||||
from pts.feature.time_feature import (DayOfMonth, DayOfWeek, DayOfYear,
|
||||
HourOfDay, MinuteOfHour, MonthOfYear,
|
||||
TimeFeature, WeekOfYear,
|
||||
time_features_from_frequency_str)
|
||||
from pts.feature.transform import (AddAgeFeature, AddConstFeature,
|
||||
AddObservedValuesIndicator, AddTimeFeatures,
|
||||
AdhocTransform, AsNumpyArray,
|
||||
CanonicalInstanceSplitter, Chain,
|
||||
ConcatFeatures, ExpandDimArray,
|
||||
FilterTransformation, FlatMapTransformation,
|
||||
IdentityTransformation, InstanceSplitter,
|
||||
ListFeatures, MapTransformation,
|
||||
RemoveFields, RenameFields, SelectFields,
|
||||
SetField, SimpleTransformation, SwapAxes,
|
||||
Transformation, VstackFeatures)
|
||||
|
||||
+2
-2
@@ -12,7 +12,7 @@
|
||||
# permissions and limitations under the License.
|
||||
|
||||
# Standard library imports
|
||||
from typing import List, Tuple, Optional
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
# Third-party imports
|
||||
import numpy as np
|
||||
@@ -122,4 +122,4 @@ def get_lags_for_frequency(freq_str: str,
|
||||
]
|
||||
lags = [1, 2, 3, 4, 5, 6, 7] + sorted(list(set(lags)))
|
||||
|
||||
return lags[:num_lags]
|
||||
return lags[:num_lags]
|
||||
|
||||
@@ -6,6 +6,7 @@ import pandas as pd
|
||||
|
||||
from .utils import get_granularity
|
||||
|
||||
|
||||
class TimeFeature(ABC):
|
||||
def __init__(self, normalized: bool = True):
|
||||
self.normalized = normalized
|
||||
|
||||
@@ -1,13 +1,14 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from collections import Counter
|
||||
from functools import lru_cache, reduce
|
||||
from typing import Iterator, List, Callable, Any, Optional, Dict, Tuple
|
||||
from typing import Any, Callable, Dict, Iterator, List, Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
|
||||
from pts.dataset import DataEntry, InstanceSampler
|
||||
from pts import assert_pts
|
||||
from pts.dataset import DataEntry, InstanceSampler
|
||||
|
||||
from .time_feature import TimeFeature
|
||||
|
||||
MAX_IDLE_TRANSFORMS = 100
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
from typing import Tuple
|
||||
import re
|
||||
from typing import Tuple
|
||||
|
||||
|
||||
def get_granularity(freq_str: str) -> Tuple[int, str]:
|
||||
"""
|
||||
@@ -18,4 +19,4 @@ def get_granularity(freq_str: str) -> Tuple[int, str]:
|
||||
groups = m.groups()
|
||||
multiple = int(groups[1]) if groups[1] is not None else 1
|
||||
granularity = groups[2]
|
||||
return multiple, granularity
|
||||
return multiple, granularity
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from pts.model.estimator import Estimator, PTSEstimator
|
||||
from pts.model.predictor import Predictor
|
||||
from pts.model.forecast import Forecast
|
||||
from pts.model.predictor import Predictor
|
||||
from pts.model.quantile import Quantile
|
||||
|
||||
@@ -1 +1 @@
|
||||
from .deepar_estimator import DeepAREstimator
|
||||
from .deepar_estimator import DeepAREstimator
|
||||
|
||||
@@ -2,9 +2,11 @@ from typing import List, Optional
|
||||
|
||||
import numpy as np
|
||||
|
||||
from pts.model import PTSEstimator
|
||||
from pts.feature import TimeFeature, get_lags_for_frequency, time_features_from_frequency_str
|
||||
from pts import Trainer
|
||||
from pts.feature import (TimeFeature, get_lags_for_frequency,
|
||||
time_features_from_frequency_str)
|
||||
from pts.model import PTSEstimator
|
||||
|
||||
|
||||
class DeepAREstimator(PTSEstimator):
|
||||
def __init__(self,
|
||||
@@ -67,4 +69,4 @@ class DeepAREstimator(PTSEstimator):
|
||||
|
||||
self.history_length = self.context_length + max(self.lags_seq)
|
||||
|
||||
self.num_parallel_samples = num_parallel_samples
|
||||
self.num_parallel_samples = num_parallel_samples
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class DeepARNetwork(nn.Module):
|
||||
pass
|
||||
pass
|
||||
|
||||
@@ -1,14 +1,13 @@
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
import numpy as np
|
||||
|
||||
from pts.dataset.common import Dataset
|
||||
from pts.dataset import TrainDataLoader
|
||||
from pts.feature import Transformation
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from pts.dataset import TrainDataLoader
|
||||
from pts.dataset.common import Dataset
|
||||
from pts.feature import Transformation
|
||||
|
||||
from .predictor import Predictor
|
||||
from .utils import get_module_forward_input_names
|
||||
|
||||
|
||||
@@ -1,9 +1,8 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Optional, Dict, Enum, Set, List
|
||||
from typing import Dict, Enum, List, Optional, Set
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
|
||||
import torch
|
||||
from torch.distributions.distribution import Distributions
|
||||
|
||||
@@ -393,4 +392,4 @@ class QuantileForecast(Forecast):
|
||||
# freq=self.freq,
|
||||
# item_id=self.item_id,
|
||||
# info=self.info,
|
||||
# )
|
||||
# )
|
||||
|
||||
@@ -2,9 +2,11 @@ from abc import ABC, abstractmethod
|
||||
from typing import Iterator
|
||||
|
||||
from pts.dataset.common import Dataset
|
||||
|
||||
from .forecast import Forecast
|
||||
from .predictor import Predictor
|
||||
|
||||
|
||||
class Predictor(ABC):
|
||||
def __init__(self, prediction_length: int, freq: str) -> None:
|
||||
self.prediction_length = prediction_length
|
||||
@@ -13,4 +15,3 @@ class Predictor(ABC):
|
||||
@abstractmethod
|
||||
def predict(self, dataset: Dataset, **kwargs) -> Iterator[Forecast]:
|
||||
pass
|
||||
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from typing import NamedTuple, Union
|
||||
import re
|
||||
from typing import NamedTuple, Union
|
||||
|
||||
|
||||
class Quantile(NamedTuple):
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
from pts.modules.distribution_output import ArgProj
|
||||
from pts.modules.lambda_layer import LambdaLayer
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from typing import Callable, Dict, Optional, Tuple
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
import numpy as np
|
||||
|
||||
@@ -8,25 +9,26 @@ import torch.nn as nn
|
||||
|
||||
class ArgProj(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_features,
|
||||
args_dim: Dict[str, int],
|
||||
domain_map: Callable[..., Tuple[torch.Tensor]],
|
||||
dtype: np.dtype = np.float32,
|
||||
prefix: Optional[str] = None,
|
||||
**kwargs,
|
||||
self,
|
||||
in_features,
|
||||
args_dim: Dict[str, int],
|
||||
domain_map: Callable[..., Tuple[torch.Tensor]],
|
||||
dtype: np.dtype = np.float32,
|
||||
prefix: Optional[str] = None,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
self.args_dim = args_dim
|
||||
self.dtype = dtype
|
||||
self.proj = nn.ModuleList([
|
||||
nn.Linear(in_features, dim)
|
||||
for dim in args_dim.values()])
|
||||
self.proj = nn.ModuleList(
|
||||
[nn.Linear(in_features, dim) for dim in args_dim.values()]
|
||||
)
|
||||
self.domain_map = domain_map
|
||||
|
||||
|
||||
def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor]:
|
||||
params_unbounded = [proj(x) for proj in self.proj]
|
||||
|
||||
return self.domain_map(*params_unbounded)
|
||||
|
||||
|
||||
class Output(ABC):
|
||||
|
||||
@@ -1,10 +1,11 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
class Lambda(nn.Module):
|
||||
|
||||
class LambdaLayer(nn.Module):
|
||||
def __init__(self, function):
|
||||
super().__init__()
|
||||
self._func = function
|
||||
|
||||
|
||||
def forward(self, x, *args):
|
||||
return self._func(x, *args)
|
||||
Reference in New Issue
Block a user