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

Class MegatronPretrainingSampler

codegeex/megatron/data/data_samplers.py:62–120  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

60
61
62class MegatronPretrainingSampler:
63 def __init__(
64 self,
65 total_samples,
66 consumed_samples,
67 micro_batch_size,
68 data_parallel_rank,
69 data_parallel_size,
70 drop_last=True,
71 ):
72 # Keep a copy of input params for later use.
73 self.total_samples = total_samples
74 self.consumed_samples = consumed_samples
75 self.micro_batch_size = micro_batch_size
76 self.data_parallel_rank = data_parallel_rank
77 self.micro_batch_times_data_parallel_size = (
78 self.micro_batch_size * data_parallel_size
79 )
80 self.drop_last = drop_last
81
82 # Sanity checks.
83 assert self.total_samples > 0, "no sample to consume: {}".format(
84 self.total_samples
85 )
86 assert (
87 self.consumed_samples < self.total_samples
88 ), "no samples left to consume: {}, {}".format(
89 self.consumed_samples, self.total_samples
90 )
91 assert self.micro_batch_size > 0
92 assert data_parallel_size > 0
93 assert (
94 self.data_parallel_rank < data_parallel_size
95 ), "data_parallel_rank should be smaller than data size: {}, " "{}".format(
96 self.data_parallel_rank, data_parallel_size
97 )
98
99 def __len__(self):
100 return self.total_samples
101
102 def get_start_end_idx(self):
103 start_idx = self.data_parallel_rank * self.micro_batch_size
104 end_idx = start_idx + self.micro_batch_size
105 return start_idx, end_idx
106
107 def __iter__(self):
108 batch = []
109 # Last batch will be dropped if drop_last is not set False
110 for idx in range(self.consumed_samples, self.total_samples):
111 batch.append(idx)
112 if len(batch) == self.micro_batch_times_data_parallel_size:
113 start_idx, end_idx = self.get_start_end_idx()
114 yield batch[start_idx:end_idx]
115 batch = []
116
117 # Check the last partial batch and see drop_last is set
118 if len(batch) > 0 and not self.drop_last:
119 start_idx, end_idx = self.get_start_end_idx()

Callers 1

Calls

no outgoing calls

Tested by

no test coverage detected