MCPcopy Create free account
hub / github.com/RingBDStack/GDAP / DistributedSortishSampler

Class DistributedSortishSampler

seq2seq/utils.py:367–417  ·  view source on GitHub ↗

Copied from torch DistributedSampler

Source from the content-addressed store, hash-verified

365
366
367class DistributedSortishSampler(Sampler):
368 """Copied from torch DistributedSampler"""
369
370 def __init__(self, dataset, batch_size, num_replicas=None, rank=None, add_extra_examples=True, shuffle=True):
371 if num_replicas is None:
372 if not dist.is_available():
373 raise RuntimeError("Requires distributed package to be available")
374 num_replicas = dist.get_world_size()
375 if rank is None:
376 if not dist.is_available():
377 raise RuntimeError("Requires distributed package to be available")
378 rank = dist.get_rank()
379 self.dataset = dataset
380 self.num_replicas = num_replicas
381 self.rank = rank
382 self.epoch = 0
383 if add_extra_examples:
384 self.num_samples = int(math.ceil(len(self.dataset) * 1.0 / self.num_replicas))
385 self.total_size = self.num_samples * self.num_replicas
386 else:
387 self.total_size = len(dataset)
388 self.num_samples = len(self.available_indices)
389 self.batch_size = batch_size
390 self.add_extra_examples = add_extra_examples
391 self.shuffle = shuffle
392
393 def __iter__(self) -> Iterable:
394 g = torch.Generator()
395 g.manual_seed(self.epoch)
396
397 sortish_data = [self.dataset.src_lens[i] for i in self.available_indices]
398 sortish_indices = sortish_sampler_indices(sortish_data, self.batch_size, shuffle=self.shuffle)
399 indices = [self.available_indices[i] for i in sortish_indices]
400 assert len(indices) == self.num_samples
401 return iter(indices)
402
403 @cached_property
404 def available_indices(self) -> np.array:
405 indices = list(range(len(self.dataset)))
406 # add extra samples to make it evenly divisible
407 indices += indices[: (self.total_size - len(indices))]
408 assert len(indices) == self.total_size
409 # subsample
410 available_indices = indices[self.rank: self.total_size: self.num_replicas]
411 return available_indices
412
413 def __len__(self):
414 return self.num_samples
415
416 def set_epoch(self, epoch):
417 self.epoch = epoch
418
419
420logger = getLogger(__name__)

Callers 1

make_sortish_samplerMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected