MCPcopy Create free account
hub / github.com/Netflix/void-model / RandomSampler

Class RandomSampler

videox_fun/data/bucket_sampler.py:50–112  ·  view source on GitHub ↗

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

Source from the content-addressed store, hash-verified

48 return img.size # (width, height)
49
50class 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

Callers 2

mainFunction · 0.90
mainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected