diff --git a/norch/utils/data/dataloader.py b/norch/utils/data/dataloader.py index 01bd41e..74a2800 100644 --- a/norch/utils/data/dataloader.py +++ b/norch/utils/data/dataloader.py @@ -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 diff --git a/norch/utils/data/distributed.py b/norch/utils/data/distributed.py new file mode 100644 index 0000000..6af54bc --- /dev/null +++ b/norch/utils/data/distributed.py @@ -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 \ No newline at end of file