MCPcopy Create free account
hub / github.com/zai-org/CodeGeeX / build_pretraining_data_loader

Function build_pretraining_data_loader

codegeex/megatron/data/data_samplers.py:24–59  ·  view source on GitHub ↗

Buld dataloader given an input dataset.

(dataset, consumed_samples)

Source from the content-addressed store, hash-verified

22
23
24def build_pretraining_data_loader(dataset, consumed_samples):
25 """Buld dataloader given an input dataset."""
26
27 if dataset is None:
28 return None
29 args = get_args()
30
31 # Megatron sampler
32 if args.dataloader_type == "single":
33 batch_sampler = MegatronPretrainingSampler(
34 total_samples=len(dataset),
35 consumed_samples=consumed_samples,
36 micro_batch_size=args.micro_batch_size,
37 data_parallel_rank=mpu.get_data_parallel_rank(),
38 data_parallel_size=mpu.get_data_parallel_world_size(),
39 )
40 elif args.dataloader_type == "cyclic":
41 batch_sampler = MegatronPretrainingRandomSampler(
42 total_samples=len(dataset),
43 consumed_samples=consumed_samples,
44 micro_batch_size=args.micro_batch_size,
45 data_parallel_rank=mpu.get_data_parallel_rank(),
46 data_parallel_size=mpu.get_data_parallel_world_size(),
47 )
48 else:
49 raise Exception(
50 "{} dataloader type is not supported.".format(args.dataloader_type)
51 )
52
53 # Torch dataloader.
54 return torch.utils.data.DataLoader(
55 dataset,
56 batch_sampler=batch_sampler,
57 num_workers=args.num_workers,
58 pin_memory=True,
59 )
60
61
62class MegatronPretrainingSampler:

Callers 1

Calls 3

get_argsFunction · 0.90

Tested by

no test coverage detected