| 180 | # for each sample in a minibatch. |
| 181 | |
| 182 | class 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. |