MCPcopy Create free account
hub / github.com/AIRMEC/HECTOR / FeatureBagsDataset

Class FeatureBagsDataset

utils.py:34–64  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

32
33
34class 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
66def define_data_sampling(train_split, val_split, method, workers):
67 # Reproducibility of DataLoader.

Callers 1

prepare_datasetsFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected