| 300 | # return eeg, label |
| 301 | |
| 302 | class 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 | |
| 328 | def create_EEG_dataset(eeg_signals_path='../dreamdiffusion/datasets/eeg_5_95_std.pth', |
no outgoing calls
no test coverage detected