(self,
dataset,
batch_size=None,
route=ROUTE_TRAIN,
pin_memory=True,
data_sampler=None,
collate_fn=None,
num_local_io_workers=None)
| 2484 | return self._step_applied |
| 2485 | |
| 2486 | def deepspeed_io(self, |
| 2487 | dataset, |
| 2488 | batch_size=None, |
| 2489 | route=ROUTE_TRAIN, |
| 2490 | pin_memory=True, |
| 2491 | data_sampler=None, |
| 2492 | collate_fn=None, |
| 2493 | num_local_io_workers=None): |
| 2494 | if not (self.is_map_style_dataset(dataset) or self.is_iterable_style_dataset(dataset)): |
| 2495 | raise ValueError("Training data must be a torch Dataset") |
| 2496 | |
| 2497 | if batch_size is None: |
| 2498 | batch_size = self.train_micro_batch_size_per_gpu() |
| 2499 | |
| 2500 | if collate_fn is None: |
| 2501 | collate_fn = self.collate_fn |
| 2502 | |
| 2503 | # Currently we only use timer in train route |
| 2504 | deepspeed_io_timer = None |
| 2505 | if route == ROUTE_TRAIN: |
| 2506 | deepspeed_io_timer = self.tput_timer |
| 2507 | |
| 2508 | # If mpu is provided, forward world size and parallel rank to sampler. |
| 2509 | data_parallel_world_size = self.dp_world_size |
| 2510 | data_parallel_rank = self.global_rank |
| 2511 | if self.mpu is not None: |
| 2512 | data_parallel_world_size = self.mpu.get_data_parallel_world_size() |
| 2513 | data_parallel_rank = self.mpu.get_data_parallel_rank() |
| 2514 | |
| 2515 | if data_sampler is None and (route == ROUTE_PREDICT or route == ROUTE_EVAL): |
| 2516 | data_sampler = torch.utils.data.DistributedSampler( |
| 2517 | dataset, |
| 2518 | num_replicas=data_parallel_world_size, |
| 2519 | rank=data_parallel_rank, |
| 2520 | shuffle=False, |
| 2521 | ) |
| 2522 | |
| 2523 | deepspeed_dataloader_config = {} |
| 2524 | if self.curriculum_learning_enabled(): |
| 2525 | deepspeed_dataloader_config = { |
| 2526 | CURRICULUM_LEARNING: self.curriculum_learning_enabled(), |
| 2527 | DATA_EFFICIENCY: self.data_efficiency_config(), |
| 2528 | DATA_PARALLEL_GROUP: self.data_parallel_group, |
| 2529 | GRADIENT_ACCUMULATION_STEPS: self.gradient_accumulation_steps(), |
| 2530 | GLOBAL_RANK: self.global_rank, |
| 2531 | DATA_SAMPLING_NUM_WORKERS: self.data_sampling_config()[DATA_SAMPLING_NUM_WORKERS] |
| 2532 | } |
| 2533 | return DeepSpeedDataLoader(dataset=dataset, |
| 2534 | batch_size=batch_size, |
| 2535 | pin_memory=pin_memory, |
| 2536 | collate_fn=collate_fn, |
| 2537 | local_rank=self.local_rank, |
| 2538 | tput_timer=deepspeed_io_timer, |
| 2539 | num_local_io_workers=num_local_io_workers, |
| 2540 | data_sampler=data_sampler, |
| 2541 | data_parallel_world_size=data_parallel_world_size, |
| 2542 | data_parallel_rank=data_parallel_rank, |
| 2543 | dataloader_drop_last=self.dataloader_drop_last(), |
no test coverage detected