| 4 | import torch |
| 5 | |
| 6 | class DiscreteSampling: |
| 7 | |
| 8 | def __init__(self, num_idx, uniform_sampling=False): |
| 9 | self.num_idx = num_idx |
| 10 | self.uniform_sampling = uniform_sampling |
| 11 | self.is_distributed = torch.distributed.is_available() and torch.distributed.is_initialized() |
| 12 | |
| 13 | # print("self.is_distributed status is ", self.is_distributed) |
| 14 | if self.is_distributed and self.uniform_sampling: |
| 15 | world_size = torch.distributed.get_world_size() |
| 16 | self.rank = torch.distributed.get_rank() |
| 17 | |
| 18 | i = 1 |
| 19 | while True: |
| 20 | if world_size % i != 0 or num_idx % (world_size // i) != 0: |
| 21 | i += 1 |
| 22 | else: |
| 23 | self.group_num = world_size // i |
| 24 | break |
| 25 | assert self.group_num > 0 |
| 26 | assert world_size % self.group_num == 0 |
| 27 | # the number of rank in one group |
| 28 | self.group_width = world_size // self.group_num |
| 29 | self.sigma_interval = self.num_idx // self.group_num |
| 30 | print('rank=%d world_size=%d group_num=%d group_width=%d sigma_interval=%s' % ( |
| 31 | self.rank, world_size, self.group_num, |
| 32 | self.group_width, self.sigma_interval)) |
| 33 | |
| 34 | |
| 35 | def __call__(self, n_samples, generator=None, device=None): |
| 36 | |
| 37 | |
| 38 | if self.is_distributed and self.uniform_sampling: |
| 39 | group_index = self.rank // self.group_width |
| 40 | idx = torch.randint( |
| 41 | group_index * self.sigma_interval, |
| 42 | (group_index + 1) * self.sigma_interval, |
| 43 | (n_samples,), |
| 44 | generator=generator, device=device, |
| 45 | ) |
| 46 | # print('proc[%d] idx=%s' % (self.rank, idx)) |
| 47 | # print("Uniform sample range is ", group_index * self.sigma_interval, (group_index + 1) * self.sigma_interval) |
| 48 | |
| 49 | else: |
| 50 | idx = torch.randint( |
| 51 | 0, self.num_idx, (n_samples,), |
| 52 | generator=generator, device=device, |
| 53 | ) |
| 54 | return idx |