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

Class EEGDataset

code/dataset.py:238–300  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

236
237
238class 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)

Callers 1

create_EEG_datasetFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected