MCPcopy Create free account
hub / github.com/NVlabs/edm / StackedRandomGenerator

Class StackedRandomGenerator

generate.py:182–196  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

180# for each sample in a minibatch.
181
182class StackedRandomGenerator:
183 def __init__(self, device, seeds):
184 super().__init__()
185 self.generators = [torch.Generator(device).manual_seed(int(seed) % (1 << 32)) for seed in seeds]
186
187 def randn(self, size, **kwargs):
188 assert size[0] == len(self.generators)
189 return torch.stack([torch.randn(size[1:], generator=gen, **kwargs) for gen in self.generators])
190
191 def randn_like(self, input):
192 return self.randn(input.shape, dtype=input.dtype, layout=input.layout, device=input.device)
193
194 def randint(self, *args, size, **kwargs):
195 assert size[0] == len(self.generators)
196 return torch.stack([torch.randint(*args, size=size[1:], generator=gen, **kwargs) for gen in self.generators])
197
198#----------------------------------------------------------------------------
199# Parse a comma separated list of numbers or ranges and return a list of ints.

Callers 1

mainFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected