MCPcopy Create free account
hub / github.com/bbaaii/DreamDiffusion / Splitter

Class Splitter

code/dataset.py:302–325  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

300 # return eeg, label
301
302class Splitter:
303
304 def __init__(self, dataset, split_path, split_num=0, split_name="train", subject=4):
305 # Set EEG dataset
306 self.dataset = dataset
307 # Load split
308 loaded = torch.load(split_path)
309
310 self.split_idx = loaded["splits"][split_num][split_name]
311 # Filter data
312 self.split_idx = [i for i in self.split_idx if i <= len(self.dataset.data) and 450 <= self.dataset.data[i]["eeg"].size(1) <= 600]
313 # Compute size
314
315 self.size = len(self.split_idx)
316 self.num_voxels = 440
317 self.data_len = 512
318
319 # Get size
320 def __len__(self):
321 return self.size
322
323 # Get item
324 def __getitem__(self, i):
325 return self.dataset[self.split_idx[i]]
326
327
328def create_EEG_dataset(eeg_signals_path='../dreamdiffusion/datasets/eeg_5_95_std.pth',

Callers 1

create_EEG_datasetFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected