Function
build_data_loader
(
data_source=None,
batch_size=64,
input_size=224,
tfm=None,
is_train=True,
shuffle=False,
dataset_wrapper=None
)
Source from the content-addressed store, hash-verified
| 352 | |
| 353 | |
| 354 | def build_data_loader( |
| 355 | data_source=None, |
| 356 | batch_size=64, |
| 357 | input_size=224, |
| 358 | tfm=None, |
| 359 | is_train=True, |
| 360 | shuffle=False, |
| 361 | dataset_wrapper=None |
| 362 | ): |
| 363 | |
| 364 | if dataset_wrapper is None: |
| 365 | dataset_wrapper = DatasetWrapper |
| 366 | |
| 367 | # Build data loader |
| 368 | data_loader = torch.utils.data.DataLoader( |
| 369 | dataset_wrapper(data_source, input_size=input_size, transform=tfm, is_train=is_train), |
| 370 | batch_size=batch_size, |
| 371 | num_workers=8, |
| 372 | shuffle=shuffle, |
| 373 | drop_last=False, |
| 374 | pin_memory=(torch.cuda.is_available()) |
| 375 | ) |
| 376 | assert len(data_loader) > 0 |
| 377 | |
| 378 | return data_loader |
Tested by
no test coverage detected