distributed sampler

This commit is contained in:
lucasdelimanogueira 2024-05-25 09:48:32 -03:00
parent a316dfd191
commit 9d7417a531
2 changed files with 42 additions and 8 deletions

View file

@ -2,19 +2,26 @@ import numpy as np
from .batch import Batch
class Dataloader(object):
def __init__(self, dataset, batch_size=32):
class Dataloader:
def __init__(self, dataset, batch_size=32, sampler=None):
self.dataset = dataset
self.batch_size = batch_size
self.sampler = sampler
def __iter__(self):
starts = np.arange(0, len(self.dataset), self.batch_size)
for start in starts:
if self.sampler is not None:
indices = list(self.sampler)
else:
indices = np.arange(len(self.dataset))
for start in range(0, len(indices), self.batch_size):
end = start + self.batch_size
batch_size = min(end, len(self.dataset)) - start
yield Batch(self.dataset[start:end], batch_size)
batch_indices = indices[start:end]
batch_size = len(batch_indices)
yield Batch([self.dataset[i] for i in batch_indices], batch_size)
def __len__(self):
if self.sampler is not None:
return len(self.sampler) // self.batch_size
return len(self.dataset) // self.batch_size

View file

@ -0,0 +1,27 @@
import math
import numpy as np
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):
# Create indices for the dataset
indices = np.arange(len(self.dataset))
# Add extra samples to make it evenly divisible
indices = np.concatenate([indices, indices[:(self.total_size - len(indices)) % 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