PyNorch/norch/utils/data/distributed.py
2024-05-27 15:08:27 -03:00

28 lines
923 B
Python

import numpy as np
import math
class DistributedSampler:
def __init__(self, dataset, num_replicas, rank):
self.dataset = dataset
self.num_replicas = num_replicas
self.rank = rank
self.num_samples = int(math.ceil(len(self.dataset) * 1.0 / self.num_replicas))
self.total_size = self.num_samples * self.num_replicas
def __iter__(self):
indices = list(range(len(self.dataset)))
# Add extra samples to make it evenly divisible
indices = indices[:self.total_size]
if len(indices) < self.total_size:
indices += indices[:(self.total_size - len(indices))]
assert len(indices) == self.total_size
# subsample
indices = indices[self.rank:self.total_size:self.num_replicas]
assert len(indices) == self.num_samples
return iter(indices)
def __len__(self):
return self.num_samples