From 6de912194040585d2116ae911fe2e254cd04d59c Mon Sep 17 00:00:00 2001 From: Kashif Rasul Date: Thu, 14 Nov 2019 21:48:37 +0100 Subject: [PATCH] fixed all but 1 test --- pts/feature/transform.py | 14 +++++++++++--- test/feature/test_transformation.py | 22 +++++++++++++++++++++- 2 files changed, 32 insertions(+), 4 deletions(-) diff --git a/pts/feature/transform.py b/pts/feature/transform.py index ac3b91c..1f4f4a1 100644 --- a/pts/feature/transform.py +++ b/pts/feature/transform.py @@ -571,8 +571,12 @@ class AddTimeFeatures(MapTransformation): self.full_date_range = pd.date_range( self._min_time_point, self._max_time_point, freq=start.freq ) - self._full_range_date_features = np.vstack( - [feat(self.full_date_range) for feat in self.date_features] + self._full_range_date_features = ( + np.vstack( + [feat(self.full_date_range) for feat in self.date_features] + ) + if self.date_features + else None ) self._date_index = pd.Series( index=self.full_date_range, data=np.arange(len(self.full_date_range)) @@ -585,7 +589,11 @@ class AddTimeFeatures(MapTransformation): ) self._update_cache(start, length) i0 = self._date_index[start] - features = self._full_range_date_features[..., i0 : i0 + length] + features = ( + self._full_range_date_features[..., i0 : i0 + length] + if self.date_features + else None + ) data[self.output_field] = features return data diff --git a/test/feature/test_transformation.py b/test/feature/test_transformation.py index ef588d0..82de788 100644 --- a/test/feature/test_transformation.py +++ b/test/feature/test_transformation.py @@ -360,7 +360,7 @@ def test_multi_dim_transformation(is_train): past_length=train_length, future_length=pred_length, time_series_fields=["dynamic_feat", "observed_values"], - output_NTC=False, + batch_first=False, ), ] ) @@ -499,3 +499,23 @@ def assert_shape(array: np.array, reference_shape: Tuple[int, int]): array.shape == reference_shape ), f"Shape should be {reference_shape} but found {array.shape}." +def assert_padded_array( + sampled_array: np.array, reference_array: np.array, padding_array: np.array +): + num_padded = int(np.sum(padding_array)) + sampled_no_padding = sampled_array[:, num_padded:] + + reference_array = np.roll(reference_array, num_padded, axis=1) + reference_no_padding = reference_array[:, num_padded:] + + # Convert nans to dummy value for assertion because + # np.nan == np.nan -> False. + reference_no_padding[np.isnan(reference_no_padding)] = 9999.0 + sampled_no_padding[np.isnan(sampled_no_padding)] = 9999.0 + + reference_no_padding = np.array(reference_no_padding, dtype=np.float32) + + assert (sampled_no_padding == reference_no_padding).all(), ( + f"Sampled and reference arrays do not match. '" + f"Got {sampled_no_padding} but should be {reference_no_padding}." + )