r"""Samples elements randomly. If without replacement, then sample from a shuffled dataset. If with replacement, then user can specify :attr:`num_samples` to draw. Args: data_source (Dataset): dataset to sample from replacement (bool): samples are drawn on-demand with repla
| 48 | return img.size # (width, height) |
| 49 | |
| 50 | class RandomSampler(Sampler[int]): |
| 51 | r"""Samples elements randomly. If without replacement, then sample from a shuffled dataset. |
| 52 | |
| 53 | If with replacement, then user can specify :attr:`num_samples` to draw. |
| 54 | |
| 55 | Args: |
| 56 | data_source (Dataset): dataset to sample from |
| 57 | replacement (bool): samples are drawn on-demand with replacement if ``True``, default=``False`` |
| 58 | num_samples (int): number of samples to draw, default=`len(dataset)`. |
| 59 | generator (Generator): Generator used in sampling. |
| 60 | """ |
| 61 | |
| 62 | data_source: Sized |
| 63 | replacement: bool |
| 64 | |
| 65 | def __init__(self, data_source: Sized, replacement: bool = False, |
| 66 | num_samples: Optional[int] = None, generator=None) -> None: |
| 67 | self.data_source = data_source |
| 68 | self.replacement = replacement |
| 69 | self._num_samples = num_samples |
| 70 | self.generator = generator |
| 71 | self._pos_start = 0 |
| 72 | |
| 73 | if not isinstance(self.replacement, bool): |
| 74 | raise TypeError(f"replacement should be a boolean value, but got replacement={self.replacement}") |
| 75 | |
| 76 | if not isinstance(self.num_samples, int) or self.num_samples <= 0: |
| 77 | raise ValueError(f"num_samples should be a positive integer value, but got num_samples={self.num_samples}") |
| 78 | |
| 79 | @property |
| 80 | def num_samples(self) -> int: |
| 81 | # dataset size might change at runtime |
| 82 | if self._num_samples is None: |
| 83 | return len(self.data_source) |
| 84 | return self._num_samples |
| 85 | |
| 86 | def __iter__(self) -> Iterator[int]: |
| 87 | n = len(self.data_source) |
| 88 | if self.generator is None: |
| 89 | seed = int(torch.empty((), dtype=torch.int64).random_().item()) |
| 90 | generator = torch.Generator() |
| 91 | generator.manual_seed(seed) |
| 92 | else: |
| 93 | generator = self.generator |
| 94 | |
| 95 | if self.replacement: |
| 96 | for _ in range(self.num_samples // 32): |
| 97 | yield from torch.randint(high=n, size=(32,), dtype=torch.int64, generator=generator).tolist() |
| 98 | yield from torch.randint(high=n, size=(self.num_samples % 32,), dtype=torch.int64, generator=generator).tolist() |
| 99 | else: |
| 100 | for _ in range(self.num_samples // n): |
| 101 | xx = torch.randperm(n, generator=generator).tolist() |
| 102 | if self._pos_start >= n: |
| 103 | self._pos_start = 0 |
| 104 | print("xx top 10", xx[:10], self._pos_start) |
| 105 | for idx in range(self._pos_start, n): |
| 106 | yield xx[idx] |
| 107 | self._pos_start = (self._pos_start + 1) % n |