| 32 | |
| 33 | |
| 34 | class FeatureBagsDataset(Dataset): |
| 35 | def __init__(self, df, data_dir, input_feature_size, stage_class): |
| 36 | self.slide_df = df.copy().reset_index(drop=True) |
| 37 | self.data_dir = data_dir |
| 38 | self.input_feature_size = input_feature_size |
| 39 | self.stage_class = stage_class |
| 40 | |
| 41 | def _get_feature_path(self, slide_id): |
| 42 | return os.path.join(self.data_dir, f"{slide_id}_Mergedfeatures.pt") |
| 43 | |
| 44 | def __getitem__(self, idx): |
| 45 | slide_id = self.slide_df["slide_id"][idx] |
| 46 | stage = self.slide_df["stage"][idx] |
| 47 | label = self.slide_df["disc_label"][idx] |
| 48 | event_time = self.slide_df["recurrence_years"][idx] |
| 49 | censorship = self.slide_df["censorship"][idx] |
| 50 | |
| 51 | full_path = self._get_feature_path(slide_id) |
| 52 | |
| 53 | features = torch.load(full_path) |
| 54 | |
| 55 | # Merged features. |
| 56 | features_merged = torch.from_numpy(np.array([x[0].mean(0) for x in features])) |
| 57 | |
| 58 | # Alternative would be all features depending on what works best. |
| 59 | features_flattened = torch.from_numpy(np.concatenate([x[0] for x in features])) |
| 60 | |
| 61 | return features_merged, features_flattened, label, event_time, censorship, stage, slide_id |
| 62 | |
| 63 | def __len__(self): |
| 64 | return len(self.slide_df) |
| 65 | |
| 66 | def define_data_sampling(train_split, val_split, method, workers): |
| 67 | # Reproducibility of DataLoader. |