MCPcopy Create free account
hub / github.com/FireRedTeam/FireRedTTS2 / BucketSampler

Class BucketSampler

bin/finetune_example/posttrain_dataloader.py:231–285  ·  view source on GitHub ↗

Groups samples of similar lengths into bins to minimize padding.

Source from the content-addressed store, hash-verified

229
230
231class BucketSampler(Sampler):
232 """
233 Groups samples of similar lengths into bins to minimize padding.
234 """
235
236 def __init__(
237 self,
238 lengths: List[int],
239 batch_size: int,
240 shuffle: bool = True,
241 is_infinite: bool = True,
242 random_seed: int = 42,
243 ):
244 self.shuffle = shuffle
245 self.batch_size = batch_size
246 self.is_infinite = is_infinite
247 self.random_seed = random_seed
248 self.local_step = 0
249 self.bins = self._create_bins(lengths, batch_size)
250
251 def _create_bins(self, lengths: List[int], batch_size: int) -> List[List[int]]:
252 indices_with_lengths = sorted(enumerate(lengths), key=lambda x: x[1])
253 bins, current_bin = [], []
254
255 for idx, _ in indices_with_lengths:
256 current_bin.append(idx)
257 if len(current_bin) >= batch_size:
258 bins.append(current_bin)
259 current_bin = []
260
261 if current_bin:
262 bins.append(current_bin)
263
264 return bins
265
266 def _shuffle_bins(self, epoch: int):
267 rng = np.random.RandomState(epoch + self.random_seed)
268 rng.shuffle(self.bins)
269 for bin_ in self.bins:
270 rng.shuffle(bin_)
271
272 def __iter__(self):
273 epoch = 0
274 while True:
275 if self.shuffle:
276 self._shuffle_bins(epoch)
277 for bin_indices in self.bins:
278 yield bin_indices
279 self.local_step += 1
280 if not self.is_infinite:
281 break
282 epoch += 1
283
284 def __len__(self):
285 return len(self.bins)
286
287
288def create_dataloaders(

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected