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

Method __call__

wan/utils/discrete_sampler.py:37–52  ·  view source on GitHub ↗
(self, n_samples, generator=None, device=None)

Source from the content-addressed store, hash-verified

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:
39 group_index = self.rank // self.group_width
40 idx = torch.randint(
41 self.start_num_idx + group_index * self.sigma_interval,
42 self.start_num_idx + (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 else:
48 idx = torch.randint(
49 self.start_num_idx, self.start_num_idx + self.num_idx, (n_samples,),
50 generator=generator, device=device,
51 )
52 return idx

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected