| 121 | |
| 122 | |
| 123 | class 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: |
no outgoing calls
no test coverage detected