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

Class MegatronPretrainingRandomSampler

codegeex/megatron/data/data_samplers.py:123–185  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

121
122
123class MegatronPretrainingRandomSampler:
124 def __init__(
125 self,
126 total_samples,
127 consumed_samples,
128 micro_batch_size,
129 data_parallel_rank,
130 data_parallel_size,
131 ):
132 # Keep a copy of input params for later use.
133 self.total_samples = total_samples
134 self.consumed_samples = consumed_samples
135 self.micro_batch_size = micro_batch_size
136 self.data_parallel_rank = data_parallel_rank
137 self.data_parallel_size = data_parallel_size
138 self.micro_batch_times_data_parallel_size = (
139 self.micro_batch_size * data_parallel_size
140 )
141 self.last_batch_size = (
142 self.total_samples % self.micro_batch_times_data_parallel_size
143 )
144
145 # Sanity checks.
146 assert self.total_samples > 0, "no sample to consume: {}".format(
147 self.total_samples
148 )
149 assert self.micro_batch_size > 0
150 assert data_parallel_size > 0
151 assert (
152 self.data_parallel_rank < data_parallel_size
153 ), "data_parallel_rank should be smaller than data size: {}, " "{}".format(
154 self.data_parallel_rank, data_parallel_size
155 )
156
157 def __len__(self):
158 return self.total_samples
159
160 def __iter__(self):
161 active_total_samples = self.total_samples - self.last_batch_size
162 self.epoch = self.consumed_samples // active_total_samples
163 current_epoch_samples = self.consumed_samples % active_total_samples
164 assert current_epoch_samples % self.micro_batch_times_data_parallel_size == 0
165
166 # data sharding and random sampling
167 bucket_size = (
168 self.total_samples // self.micro_batch_times_data_parallel_size
169 ) * self.micro_batch_size
170 bucket_offset = current_epoch_samples // self.data_parallel_size
171 start_idx = self.data_parallel_rank * bucket_size
172
173 g = torch.Generator()
174 g.manual_seed(self.epoch)
175 random_idx = torch.randperm(bucket_size, generator=g).tolist()
176 idx_range = [start_idx + x for x in random_idx[bucket_offset:]]
177
178 batch = []
179 # Last batch if not complete will be dropped.
180 for idx in idx_range:

Callers 1

Calls

no outgoing calls

Tested by

no test coverage detected