(self, n_samples, generator=None, device=None)
| 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 |
nothing calls this directly
no outgoing calls
no test coverage detected