diff --git a/v1/src/timesfm/data_loader.py b/v1/src/timesfm/data_loader.py index d81b130..e2fc259 100644 --- a/v1/src/timesfm/data_loader.py +++ b/v1/src/timesfm/data_loader.py @@ -149,11 +149,12 @@ class TimeSeriesdata(object): else: epoch_len = self.epoch_len for idx in perm[0:epoch_len]: - for _ in range(num_ts // self.batch_size + 1): + batch_indices = range(0, num_ts, self.batch_size) + for batch_idx in batch_indices: if self.permute: tsidx = np.random.choice(num_ts, size=self.batch_size, replace=False) else: - tsidx = np.arange(num_ts) + tsidx = np.arange(batch_idx, min(batch_idx + self.batch_size, num_ts)) dtimes = np.arange(idx - hist_len, idx + self.pred_len) ( bts_train, diff --git a/v1/tests/test_data_loader.py b/v1/tests/test_data_loader.py new file mode 100644 index 0000000..ffee145 --- /dev/null +++ b/v1/tests/test_data_loader.py @@ -0,0 +1,50 @@ +from pathlib import Path + +import numpy as np +import pandas as pd + +from timesfm.data_loader import TimeSeriesdata + + +def test_train_gen_respects_batch_size_when_permute_is_false(tmp_path: Path) -> None: + rows = 12 + df = pd.DataFrame( + { + "ds": pd.date_range("2024-01-01", periods=rows, freq="D"), + "ts_1": np.arange(rows), + "ts_2": np.arange(rows) + 10, + "ts_3": np.arange(rows) + 20, + "ts_4": np.arange(rows) + 30, + "ts_5": np.arange(rows) + 40, + } + ) + data_path = tmp_path / "sample.csv" + df.to_csv(data_path, index=False) + + loader = TimeSeriesdata( + data_path=str(data_path), + datetime_col="ds", + num_cov_cols=None, + cat_cov_cols=None, + ts_cols=np.array(["ts_1", "ts_2", "ts_3", "ts_4", "ts_5"]), + train_range=[0, 8], + val_range=[8, 10], + test_range=[10, 12], + hist_len=3, + pred_len=2, + batch_size=2, + freq="D", + normalize=False, + epoch_len=1, + holiday=False, + permute=False, + ) + + batches = list(loader.train_gen()) + ts_indices = [batch[-1].tolist() for batch in batches] + + assert ts_indices == [[0, 1], [2, 3], [4]] + for batch in batches: + assert len(batch[-1]) <= 2 + assert batch[0].shape[0] == len(batch[-1]) + assert batch[3].shape[0] == len(batch[-1])