MCPcopy Create free account
hub / github.com/AIS-SNU/Smart-Infinity / make_data_loader

Function make_data_loader

DeepSpeedExample/megatron/utils.py:92–115  ·  view source on GitHub ↗

Buld dataloader given an input dataset.

(dataset)

Source from the content-addressed store, hash-verified

90
91
92def make_data_loader(dataset):
93 """Buld dataloader given an input dataset."""
94 if dataset is None:
95 return None
96 args = get_args()
97
98 # Data parallel arguments.
99 world_size = mpu.get_data_parallel_world_size()
100 rank = mpu.get_data_parallel_rank()
101 global_batch_size = args.batch_size * world_size
102 num_workers = args.num_workers
103
104 # Use a simple sampler with distributed batch sampler.
105 sampler = torch.utils.data.SequentialSampler(dataset)
106 batch_sampler = DistributedBatchSampler(sampler=sampler,
107 batch_size=global_batch_size,
108 drop_last=True,
109 rank=rank,
110 world_size=world_size)
111 # Torch dataloader.
112 return torch.utils.data.DataLoader(dataset,
113 batch_sampler=batch_sampler,
114 num_workers=num_workers,
115 pin_memory=True)
116
117def average_losses_across_data_parallel_group(losses):
118 """Reduce a tensor of losses across all GPUs."""

Callers 1

Calls 4

get_argsFunction · 0.90

Tested by

no test coverage detected