MCPcopy Create free account
hub / github.com/Francis-Rings/FlashPortrait / __init__

Method __init__

wan/utils/discrete_sampler.py:6–35  ·  view source on GitHub ↗
(self, num_idx, uniform_sampling=False, start_num_idx=0, sp_size=1)

Source from the content-addressed store, hash-verified

4
5class DiscreteSampling:
6 def __init__(self, num_idx, uniform_sampling=False, start_num_idx=0, sp_size=1):
7 self.num_idx = num_idx
8 self.start_num_idx = start_num_idx
9 self.uniform_sampling = uniform_sampling
10 self.is_distributed = torch.distributed.is_available() and torch.distributed.is_initialized()
11
12 if self.is_distributed and self.uniform_sampling:
13 world_size = torch.distributed.get_world_size()
14 self.rank = torch.distributed.get_rank()
15
16 i = 1
17 while True:
18 if world_size % i != 0 or num_idx % (world_size // i) != 0:
19 i += 1
20 else:
21 if i >= sp_size:
22 self.group_num = world_size // i
23 elif sp_size > world_size:
24 self.group_num = 1
25 else:
26 self.group_num = world_size // sp_size
27 break
28 assert self.group_num > 0
29 assert world_size % self.group_num == 0
30 # the number of rank in one group
31 self.group_width = world_size // self.group_num
32 self.sigma_interval = self.num_idx // self.group_num
33 print('rank=%d world_size=%d group_num=%d group_width=%d sigma_interval=%s' % (
34 self.rank, world_size, self.group_num,
35 self.group_width, self.sigma_interval))
36
37 def __call__(self, n_samples, generator=None, device=None):
38 if self.is_distributed and self.uniform_sampling:

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected