Merge pull request #97 from google-research/rajat_dev
normalization changes to data loader
This commit is contained in:
@@ -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,18 +91,16 @@ 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
|
||||||
self.freq = freq
|
self.freq = freq
|
||||||
self.normalize = normalize
|
self.normalize = normalize
|
||||||
self.data_mat = self.data_df[self.ts_cols].to_numpy().transpose()
|
self.data_mat = self.data_df[self.ts_cols].to_numpy().transpose()
|
||||||
self.data_mat = self.data_mat[:, 0 : self.test_range[1]]
|
self.data_mat = self.data_mat[:, 0:self.test_range[1]]
|
||||||
self.time_mat = self.time_df.to_numpy().transpose()
|
self.time_mat = self.time_df.to_numpy().transpose()
|
||||||
self.num_feat_mat = self.data_df[num_cov_cols].to_numpy().transpose()
|
self.num_feat_mat = self.data_df[num_cov_cols].to_numpy().transpose()
|
||||||
self.cat_feat_mat, self.cat_sizes = self._get_cat_cols(cat_cov_cols)
|
self.cat_feat_mat, self.cat_sizes = self._get_cat_cols(cat_cov_cols)
|
||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user