ran isort

isort --recursive --atomic --apply pts
This commit is contained in:
Kashif Rasul
2019-10-30 09:38:19 +01:00
parent 52ed9d8937
commit 4ad01ea2e3
25 changed files with 79 additions and 97 deletions
+1 -1
View File
@@ -1,2 +1,2 @@
from .exception import assert_pts
from .trainer import Trainer
from .trainer import Trainer
+5 -9
View File
@@ -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 -2
View File
@@ -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 -1
View File
@@ -1,6 +1,6 @@
from typing import Iterable
from .common import Dataset, DataEntry, SourceContext
from .common import DataEntry, Dataset, SourceContext
from .process import ProcessDataEntry
+3 -2
View File
@@ -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 -1
View File
@@ -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
View File
@@ -1,7 +1,6 @@
import math
from collections import defaultdict
from typing import Optional
import math
import numpy as np
+4 -3
View File
@@ -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):
+1 -1
View File
@@ -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
View File
@@ -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
View File
@@ -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]
+1
View File
@@ -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
+3 -2
View File
@@ -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
+3 -2
View File
@@ -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 -1
View File
@@ -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
View File
@@ -1 +1 @@
from .deepar_estimator import DeepAREstimator
from .deepar_estimator import DeepAREstimator
+5 -3
View File
@@ -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
+2 -1
View File
@@ -1,5 +1,6 @@
import torch
import torch.nn as nn
class DeepARNetwork(nn.Module):
pass
pass
+4 -5
View File
@@ -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
+2 -3
View File
@@ -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 -1
View File
@@ -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 -1
View File
@@ -1,5 +1,5 @@
from typing import NamedTuple, Union
import re
from typing import NamedTuple, Union
class Quantile(NamedTuple):
+2
View File
@@ -0,0 +1,2 @@
from pts.modules.distribution_output import ArgProj
from pts.modules.lambda_layer import LambdaLayer
+14 -12
View File
@@ -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)