| 107 | |
| 108 | |
| 109 | class DistributedEpochSampler(Sampler): |
| 110 | def __init__(self, dataset, total_iter, batch_size, world_size=None, rank=None, last_iter=0): |
| 111 | if world_size is None: |
| 112 | world_size = link.get_world_size() |
| 113 | if rank is None: |
| 114 | rank = link.get_rank() |
| 115 | assert rank < world_size |
| 116 | self.dataset = dataset |
| 117 | self.total_iter = total_iter |
| 118 | self.batch_size = batch_size |
| 119 | self.world_size = world_size |
| 120 | self.rank = rank |
| 121 | self.last_iter = last_iter |
| 122 | |
| 123 | self.all_size_single = self.total_iter * self.batch_size |
| 124 | |
| 125 | self.indices = self.gen_new_list() |
| 126 | self.call = 0 |
| 127 | |
| 128 | def __iter__(self): |
| 129 | if self.call == 0: |
| 130 | self.call = 1 |
| 131 | return iter(self.indices[self.last_iter*self.batch_size:]) |
| 132 | else: |
| 133 | raise RuntimeError( |
| 134 | "this sampler is not designed to be called more than once!!") |
| 135 | |
| 136 | def get_one_epoch_self_part(self): |
| 137 | num = len(self.dataset) |
| 138 | indices = np.arange(num) |
| 139 | extra_indices = np.random.choice( |
| 140 | num, self.extra_per_epoch, replace=False) |
| 141 | indices = np.concatenate((indices, extra_indices)) |
| 142 | np.random.shuffle(indices) |
| 143 | assert len(indices) % (self.world_size * self.batch_size) == 0 |
| 144 | num_single = len(indices) // self.world_size |
| 145 | return indices[self.rank*num_single:(self.rank+1)*num_single] |
| 146 | |
| 147 | def gen_new_list(self): |
| 148 | np.random.seed(0) |
| 149 | |
| 150 | self.all_num = self.total_iter * self.batch_size * self.world_size |
| 151 | iter_per_epoch = (len(self.dataset) - |
| 152 | 1) // (self.batch_size * self.world_size) + 1 |
| 153 | self.num_per_epoch = iter_per_epoch * self.batch_size * self.world_size |
| 154 | self.extra_per_epoch = self.num_per_epoch - len(self.dataset) |
| 155 | repeat = (self.all_num - 1) // self.num_per_epoch + 1 |
| 156 | indices = [] |
| 157 | for i in range(repeat): |
| 158 | indice = self.get_one_epoch_self_part() |
| 159 | indices.append(indice) |
| 160 | |
| 161 | indices = np.concatenate(indices) |
| 162 | indices = indices[:self.all_size_single] |
| 163 | |
| 164 | assert len(indices) == self.all_size_single |
| 165 | |
| 166 | return indices |
nothing calls this directly
no outgoing calls
no test coverage detected