(seed)
| 355 | dataset_library = __import__(dataset_info[0], fromlist=[dataset_info[1]]) |
| 356 | |
| 357 | def get_dataloaders(seed): |
| 358 | dataloaders = [] |
| 359 | for subdataset in subdatasets: |
| 360 | train_dataset = dataset_library.__dict__[dataset_info[1]]( |
| 361 | data_path, |
| 362 | classname=subdataset, |
| 363 | resize=resize, |
| 364 | train_val_split=train_val_split, |
| 365 | imagesize=imagesize, |
| 366 | split=dataset_library.DatasetSplit.TRAIN, |
| 367 | seed=seed, |
| 368 | augment=augment, |
| 369 | ) |
| 370 | |
| 371 | test_dataset = dataset_library.__dict__[dataset_info[1]]( |
| 372 | data_path, |
| 373 | classname=subdataset, |
| 374 | resize=resize, |
| 375 | imagesize=imagesize, |
| 376 | split=dataset_library.DatasetSplit.TEST, |
| 377 | seed=seed, |
| 378 | ) |
| 379 | |
| 380 | train_dataloader = torch.utils.data.DataLoader( |
| 381 | train_dataset, |
| 382 | batch_size=batch_size, |
| 383 | shuffle=False, |
| 384 | num_workers=num_workers, |
| 385 | pin_memory=True, |
| 386 | ) |
| 387 | |
| 388 | test_dataloader = torch.utils.data.DataLoader( |
| 389 | test_dataset, |
| 390 | batch_size=batch_size, |
| 391 | shuffle=False, |
| 392 | num_workers=num_workers, |
| 393 | pin_memory=True, |
| 394 | ) |
| 395 | |
| 396 | train_dataloader.name = name |
| 397 | if subdataset is not None: |
| 398 | train_dataloader.name += "_" + subdataset |
| 399 | |
| 400 | if train_val_split < 1: |
| 401 | val_dataset = dataset_library.__dict__[dataset_info[1]]( |
| 402 | data_path, |
| 403 | classname=subdataset, |
| 404 | resize=resize, |
| 405 | train_val_split=train_val_split, |
| 406 | imagesize=imagesize, |
| 407 | split=dataset_library.DatasetSplit.VAL, |
| 408 | seed=seed, |
| 409 | ) |
| 410 | |
| 411 | val_dataloader = torch.utils.data.DataLoader( |
| 412 | val_dataset, |
| 413 | batch_size=batch_size, |
| 414 | shuffle=False, |
nothing calls this directly
no outgoing calls
no test coverage detected