rename to batch_first

This commit is contained in:
Kashif Rasul
2019-07-18 15:37:23 +02:00
parent 3a9f945f25
commit c0f84f6a02
+9 -9
View File
@@ -667,7 +667,7 @@ class InstanceSplitter(FlatMapTransformation):
length of the target seen before making prediction
future_length
length of the target that must be predicted
output_NTC
batch_first
whether to have time series output in (time, dimension) or in
(dimension, time) layout
time_series_fields
@@ -690,7 +690,7 @@ class InstanceSplitter(FlatMapTransformation):
train_sampler: InstanceSampler,
past_length: int,
future_length: int,
output_NTC: bool = True,
batch_first: bool = True,
time_series_fields: Optional[List[str]] = None,
pick_incomplete: bool = True,
) -> None:
@@ -700,7 +700,7 @@ class InstanceSplitter(FlatMapTransformation):
self.train_sampler = train_sampler
self.past_length = past_length
self.future_length = future_length
self.output_NTC = output_NTC
self.batch_first = batch_first
self.ts_fields = time_series_fields if time_series_fields is not None else []
self.target_field = target_field
self.is_pad_field = is_pad_field
@@ -764,7 +764,7 @@ class InstanceSplitter(FlatMapTransformation):
if pad_length > 0:
pad_indicator[:pad_length] = 1
if self.output_NTC:
if self.batch_first:
for ts_field in slice_cols:
d[self._past(ts_field)] = d[self._past(ts_field)].transpose()
d[self._future(ts_field)] = d[self._future(ts_field)].transpose()
@@ -807,7 +807,7 @@ class CanonicalInstanceSplitter(FlatMapTransformation):
instance sampler that provides sampling indices given a time-series
instance_length
length of the target seen before making prediction
output_NTC
batch_first
whether to have time series output in (time, dimension) or in
(dimension, time) layout
time_series_fields
@@ -832,7 +832,7 @@ class CanonicalInstanceSplitter(FlatMapTransformation):
forecast_start_field: str,
instance_sampler: InstanceSampler,
instance_length: int,
output_NTC: bool = True,
batch_first: bool = True,
time_series_fields: List[str] = [],
allow_target_padding: bool = False,
pad_value: float = 0.0,
@@ -841,7 +841,7 @@ class CanonicalInstanceSplitter(FlatMapTransformation):
) -> None:
self.instance_sampler = instance_sampler
self.instance_length = instance_length
self.output_NTC = output_NTC
self.batch_first = batch_first
self.dynamic_feature_fields = time_series_fields
self.target_field = target_field
self.allow_target_padding = allow_target_padding
@@ -911,14 +911,14 @@ class CanonicalInstanceSplitter(FlatMapTransformation):
else:
past_ts = full_ts[..., (i - self.instance_length): i]
past_ts = past_ts.transpose() if self.output_NTC else past_ts
past_ts = past_ts.transpose() if self.batch_first else past_ts
d[self._past(ts_field)] = past_ts
if self.use_prediction_features and not is_train:
if not ts_field == self.target_field:
future_ts = full_ts[..., i: i + self.prediction_length]
future_ts = (
future_ts.transpose() if self.output_NTC else future_ts
future_ts.transpose() if self.batch_first else future_ts
)
d[self._future(ts_field)] = future_ts