Groups samples of similar lengths into bins to minimize padding.
| 229 | |
| 230 | |
| 231 | class 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 | |
| 288 | def create_dataloaders( |
nothing calls this directly
no outgoing calls
no test coverage detected