MCPcopy Create free account
hub / github.com/deepspeedai/DeepSpeed / deepspeed_io

Method deepspeed_io

deepspeed/runtime/engine.py:2486–2544  ·  view source on GitHub ↗
(self,
                     dataset,
                     batch_size=None,
                     route=ROUTE_TRAIN,
                     pin_memory=True,
                     data_sampler=None,
                     collate_fn=None,
                     num_local_io_workers=None)

Source from the content-addressed store, hash-verified

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(),

Callers 2

__init__Method · 0.95
_build_data_iterMethod · 0.80

Tested by

no test coverage detected