MCPcopy Create free account
hub / github.com/UVA-Computer-Vision-Lab/FrameINO / DiscreteSampling

Class DiscreteSampling

architecture/noise_sampler.py:6–54  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

4import torch
5
6class 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

Callers 2

mainFunction · 0.90
mainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected