From b2ce697abe4482083110e2b6e25375fd58b26e0f Mon Sep 17 00:00:00 2001 From: Lars Bogner Date: Fri, 21 Nov 2025 09:24:14 +0100 Subject: [PATCH] Allow saving of datasets --- Code/python/src/aiRNN/dataloader.py | 46 +++++++++++++++++++++++++++++ 1 file changed, 46 insertions(+) diff --git a/Code/python/src/aiRNN/dataloader.py b/Code/python/src/aiRNN/dataloader.py index 4993e5e..5a7019a 100644 --- a/Code/python/src/aiRNN/dataloader.py +++ b/Code/python/src/aiRNN/dataloader.py @@ -102,7 +102,53 @@ class BaseDataset(Dataset): df.loc[:, "last_timestamp"] = last_row["timestamp"] return df + + def save_entire_dataset(self, filepath): + """Utility to save the entire dataset to a single pth file.""" + X_feat_list = [] + X_time_list = [] + Y_out_list = [] + X_context_list = [] + for i in range(len(self)): + X_f, X_t, Y, X_c = self[i] + X_feat_list.append(X_f.cpu()) + X_time_list.append(X_t.cpu()) + Y_out_list.append(Y.cpu()) + if X_c is not None: + X_context_list.append(X_c.cpu()) + data_dict = { + "X_feat": torch.stack(X_feat_list), + "X_time": torch.stack(X_time_list), + "Y_out": torch.stack(Y_out_list), + } + if len(X_context_list) > 0: + data_dict["X_context"] = torch.stack(X_context_list) + torch.save(data_dict, filepath) +class SaveDataset(Dataset): + def __init__(self, data_dict, device="cpu"): + super().__init__() + self.X_feat = data_dict["X_feat"].to(device) + self.X_time = data_dict["X_time"].to(device) + self.Y_out = data_dict["Y_out"].to(device) + self.device = device + if "X_context" in data_dict: + self.X_context = data_dict["X_context"].to(device) + else: + self.X_context = None + + def __len__(self): + return self.X_feat.shape[0] + + def __getitem__(self, idx): + X_f = self.X_feat[idx] + X_t = self.X_time[idx] + Y = self.Y_out[idx] + if self.X_context is not None: + X_c = self.X_context[idx] + else: + X_c = None + return X_f, X_t, Y, X_c class EvenlySpacedDataset(BaseDataset): def __init__(