normalization changes to data loader

This commit is contained in:
Rajat Sen
2024-07-12 16:11:16 +00:00
parent 03c71634e6
commit 57d5cd08f7
+7 -13
View File
@@ -11,13 +11,11 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
"""TF dataloaders for general timeseries datasets. """TF dataloaders for general timeseries datasets.
The expected input format is csv file with a datetime index. The expected input format is csv file with a datetime index.
""" """
from absl import logging from absl import logging
import numpy as np import numpy as np
import pandas as pd import pandas as pd
@@ -79,9 +77,8 @@ class TimeSeriesdata(object):
self.data_df['ccol'] = np.zeros(self.data_df.shape[0]) self.data_df['ccol'] = np.zeros(self.data_df.shape[0])
cat_cov_cols = ['ccol'] cat_cov_cols = ['ccol']
self.data_df.fillna(0, inplace=True) self.data_df.fillna(0, inplace=True)
self.data_df.set_index( self.data_df.set_index(pd.DatetimeIndex(self.data_df[datetime_col]),
pd.DatetimeIndex(self.data_df[datetime_col]), inplace=True inplace=True)
)
self.num_cov_cols = num_cov_cols self.num_cov_cols = num_cov_cols
self.cat_cov_cols = cat_cov_cols self.cat_cov_cols = cat_cov_cols
self.ts_cols = ts_cols self.ts_cols = ts_cols
@@ -94,11 +91,9 @@ class TimeSeriesdata(object):
data_df_idx[-1] + pd.Timedelta(1, freq=freq), data_df_idx[-1] + pd.Timedelta(1, freq=freq),
periods=pred_len + 1, periods=pred_len + 1,
freq=freq, freq=freq,
) ))
)
self.time_df = time_features.TimeCovariates( self.time_df = time_features.TimeCovariates(
date_index, holiday=holiday date_index, holiday=holiday).get_covariates()
).get_covariates()
self.hist_len = hist_len self.hist_len = hist_len
self.pred_len = pred_len self.pred_len = pred_len
self.batch_size = batch_size self.batch_size = batch_size
@@ -135,7 +130,7 @@ class TimeSeriesdata(object):
def _normalize_data(self): def _normalize_data(self):
self.scaler = StandardScaler() self.scaler = StandardScaler()
train_mat = self.data_mat[:, self.train_range[0] : self.train_range[1]] train_mat = self.data_mat[:, 0:self.train_range[1]]
self.scaler = self.scaler.fit(train_mat.transpose()) self.scaler = self.scaler.fit(train_mat.transpose())
self.data_mat = self.scaler.transform(self.data_mat.transpose()).transpose() self.data_mat = self.scaler.transform(self.data_mat.transpose()).transpose()
@@ -253,9 +248,8 @@ class TimeSeriesdata(object):
gen_fn = self.train_gen gen_fn = self.train_gen
else: else:
gen_fn = lambda: self.test_val_gen(mode, shift) gen_fn = lambda: self.test_val_gen(mode, shift)
output_types = tuple( output_types = tuple([tf.float32] * 2 + [tf.int32] + [tf.float32] * 2 +
[tf.float32] * 2 + [tf.int32] + [tf.float32] * 2 + [tf.int32] * 2 [tf.int32] * 2)
)
dataset = tf.data.Dataset.from_generator(gen_fn, output_types) dataset = tf.data.Dataset.from_generator(gen_fn, output_types)
dataset = dataset.prefetch(tf.data.experimental.AUTOTUNE) dataset = dataset.prefetch(tf.data.experimental.AUTOTUNE)
return dataset return dataset