| 236 | |
| 237 | |
| 238 | class EEGDataset(Dataset): |
| 239 | |
| 240 | # Constructor |
| 241 | def __init__(self, eeg_signals_path, imagenet_path, image_transform=identity, subject = 4): |
| 242 | # Load EEG signals |
| 243 | loaded = torch.load(eeg_signals_path) |
| 244 | # if opt.subject!=0: |
| 245 | # self.data = [loaded['dataset'][i] for i in range(len(loaded['dataset']) ) if loaded['dataset'][i]['subject']==opt.subject] |
| 246 | # else: |
| 247 | # print(loaded) |
| 248 | if subject!=0: |
| 249 | self.data = [loaded['dataset'][i] for i in range(len(loaded['dataset']) ) if loaded['dataset'][i]['subject']==subject] |
| 250 | else: |
| 251 | self.data = loaded['dataset'] |
| 252 | self.labels = loaded["labels"] |
| 253 | self.images = loaded["images"] |
| 254 | self.imagenet = imagenet_path |
| 255 | self.image_transform = image_transform |
| 256 | self.num_voxels = 440 |
| 257 | self.data_len = 512 |
| 258 | # Compute size |
| 259 | self.size = len(self.data) |
| 260 | self.processor = AutoProcessor.from_pretrained("openai/clip-vit-large-patch14") |
| 261 | |
| 262 | # Get size |
| 263 | def __len__(self): |
| 264 | return self.size |
| 265 | |
| 266 | # Get item |
| 267 | def __getitem__(self, i): |
| 268 | |
| 269 | eeg = self.data[i]["eeg"].float().t() |
| 270 | |
| 271 | eeg = eeg[20:460,:] |
| 272 | |
| 273 | eeg = np.array(eeg.transpose(0,1)) |
| 274 | x = np.linspace(0, 1, eeg.shape[-1]) |
| 275 | x2 = np.linspace(0, 1, self.data_len) |
| 276 | f = interp1d(x, eeg) |
| 277 | eeg = f(x2) |
| 278 | eeg = torch.from_numpy(eeg).float() |
| 279 | |
| 280 | label = torch.tensor(self.data[i]["label"]).long() |
| 281 | |
| 282 | # Get label |
| 283 | image_name = self.images[self.data[i]["image"]] |
| 284 | if self.imagenet: |
| 285 | image_path = os.path.join(self.imagenet, image_name.split('_')[0], image_name+'.JPEG') |
| 286 | image_raw = Image.open(image_path).convert('RGB') |
| 287 | # print(image_path) |
| 288 | else: |
| 289 | noise = np.random.randint(0, 256, (512, 512, 3), dtype=np.uint8) |
| 290 | image_raw = Image.fromarray(noise, 'RGB') |
| 291 | |
| 292 | |
| 293 | image = np.array(image_raw) / 255.0 |
| 294 | image_raw = self.processor(images=image_raw, return_tensors="pt") |
| 295 | image_raw['pixel_values'] = image_raw['pixel_values'].squeeze(0) |
no outgoing calls
no test coverage detected