| 116 | |
| 117 | |
| 118 | class DistributedKRepeatSampler(Sampler): |
| 119 | def __init__(self, dataset, batch_size, k, num_replicas, rank, seed=0): |
| 120 | self.dataset = dataset |
| 121 | self.batch_size = batch_size |
| 122 | self.k = k |
| 123 | self.num_replicas = num_replicas |
| 124 | self.rank = rank |
| 125 | self.seed = seed |
| 126 | |
| 127 | self.total_samples = self.num_replicas * self.batch_size |
| 128 | assert ( |
| 129 | self.total_samples % self.k == 0 |
| 130 | ), f"k can not div n*b, k{k}-num_replicas{num_replicas}-batch_size{batch_size}" |
| 131 | self.m = self.total_samples // self.k |
| 132 | self.epoch = 0 |
| 133 | |
| 134 | def __iter__(self): |
| 135 | while True: |
| 136 | g = torch.Generator() |
| 137 | g.manual_seed(self.seed + self.epoch) |
| 138 | indices = torch.randperm(len(self.dataset), generator=g)[: self.m].tolist() |
| 139 | repeated_indices = [idx for idx in indices for _ in range(self.k)] |
| 140 | |
| 141 | shuffled_indices = torch.randperm(len(repeated_indices), generator=g).tolist() |
| 142 | shuffled_samples = [repeated_indices[i] for i in shuffled_indices] |
| 143 | |
| 144 | per_card_samples = [] |
| 145 | for i in range(self.num_replicas): |
| 146 | start = i * self.batch_size |
| 147 | end = start + self.batch_size |
| 148 | per_card_samples.append(shuffled_samples[start:end]) |
| 149 | yield per_card_samples[self.rank] |
| 150 | |
| 151 | def set_epoch(self, epoch): |
| 152 | self.epoch = epoch |
| 153 | |
| 154 | |
| 155 | def gather_tensor_to_all(tensor, world_size): |