| 150 | return self._num_samples |
| 151 | |
| 152 | def __iter__(self) -> Iterator[int]: |
| 153 | n = len(self.data_source) |
| 154 | if self.generator is None: |
| 155 | seed = int(torch.empty((), dtype=torch.int64).random_().item()) |
| 156 | generator = torch.Generator() |
| 157 | generator.manual_seed(seed) |
| 158 | else: |
| 159 | generator = self.generator |
| 160 | |
| 161 | if self.replacement: |
| 162 | for _ in range(self.num_samples // 32): |
| 163 | yield from torch.randint(high=n, size=(32,), dtype=torch.int64, generator=generator).tolist() |
| 164 | yield from torch.randint(high=n, size=(self.num_samples % 32,), dtype=torch.int64, generator=generator).tolist() |
| 165 | else: |
| 166 | for _ in range(self.num_samples // n): |
| 167 | yield from torch.randperm(n, generator=generator).tolist() |
| 168 | yield from torch.randperm(n, generator=generator).tolist()[:self.num_samples % n] |
| 169 | |
| 170 | def __len__(self) -> int: |
| 171 | return self.num_samples |