distributed sampler
This commit is contained in:
parent
a316dfd191
commit
9d7417a531
2 changed files with 42 additions and 8 deletions
|
|
@ -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
|
||||
|
|
|
|||
27
norch/utils/data/distributed.py
Normal file
27
norch/utils/data/distributed.py
Normal 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
|
||||
Loading…
Add table
Add a link
Reference in a new issue