PyNorch/norch/utils/data/distributed.py

29 lines
923 B
Python
Raw Permalink Normal View History

2024-05-25 09:48:32 -03:00
import numpy as np
2024-05-25 10:53:09 -03:00
import math
2024-05-25 09:48:32 -03:00
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):
2024-05-25 10:53:09 -03:00
indices = list(range(len(self.dataset)))
2024-05-25 09:48:32 -03:00
# Add extra samples to make it evenly divisible
2024-05-25 10:53:09 -03:00
indices = indices[:self.total_size]
2024-05-27 15:08:27 -03:00
if len(indices) < self.total_size:
indices += indices[:(self.total_size - len(indices))]
2024-05-25 09:48:32 -03:00
assert len(indices) == self.total_size
2024-05-25 10:53:09 -03:00
# subsample
2024-05-25 09:48:32 -03:00
indices = indices[self.rank:self.total_size:self.num_replicas]
assert len(indices) == self.num_samples
2024-05-25 10:53:09 -03:00
2024-05-25 09:48:32 -03:00
return iter(indices)
def __len__(self):
2024-05-25 10:53:09 -03:00
return self.num_samples