(
self,
batch_size: int,
num_workers: int,
dataset_config: dict[str, Any],
worker_init: bool = True,
pre_collation_factor: int = 1,
)
| 365 | assert self.batch_size % self.pre_collation_factor == 0 |
| 366 | |
| 367 | def make_loader( |
| 368 | self, |
| 369 | batch_size: int, |
| 370 | num_workers: int, |
| 371 | dataset_config: dict[str, Any], |
| 372 | worker_init: bool = True, |
| 373 | pre_collation_factor: int = 1, |
| 374 | ): |
| 375 | loader = torch.utils.data.DataLoader( |
| 376 | BilliardSimDataset(**dataset_config), |
| 377 | batch_size=batch_size // pre_collation_factor, |
| 378 | num_workers=num_workers, |
| 379 | worker_init_fn=worker_init_fn if worker_init else None, |
| 380 | collate_fn=dict_collation_fn, |
| 381 | pin_memory=True, |
| 382 | ) |
| 383 | if pre_collation_factor > 1: |
| 384 | loader = CollatingDataLoader(loader, pre_collation_factor, dict_join_collation_fn) |
| 385 | return loader |
| 386 | |
| 387 | def train_dataloader(self): |
| 388 | return self.make_loader( |
no test coverage detected