| 60 | |
| 61 | |
| 62 | class 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() |
no outgoing calls
no test coverage detected