From e3bb882d3c06b20b7b347382ad0b791002f93b86 Mon Sep 17 00:00:00 2001 From: lucasdelimanogueira Date: Mon, 27 May 2024 15:08:27 -0300 Subject: [PATCH] fix distributed dataloader --- norch/utils/data/dataloader.py | 21 +++++----- norch/utils/data/distributed.py | 2 + tests/test_distributed.py | 71 ++++++++++++++++++++++++++------- train.py | 2 +- 4 files changed, 69 insertions(+), 27 deletions(-) diff --git a/norch/utils/data/dataloader.py b/norch/utils/data/dataloader.py index 8fb8142..fa823ad 100644 --- a/norch/utils/data/dataloader.py +++ b/norch/utils/data/dataloader.py @@ -4,24 +4,23 @@ from .batch import Batch class DataLoader: - def __init__(self, dataset, batch_size=32, sampler=None): + def __init__(self, dataset, batch_size, sampler=None): self.dataset = dataset self.batch_size = batch_size self.sampler = sampler def __iter__(self): if self.sampler is not None: - indices = iter(self.sampler) - + indices = list(iter(self.sampler)) + else: - indices = range(len(self.dataset)) - - for idx in indices: - start = idx * self.batch_size - end = min(start + self.batch_size, len(self.dataset)) - yield Batch(self.dataset[start:end], end - start) - - + indices = list(range(len(self.dataset))) + + for i in range(0, len(indices), self.batch_size): + batch_indices = indices[i:i + self.batch_size] + batch_data = self.dataset[batch_indices] + yield Batch(batch_data, len(batch_data)) + def __len__(self): if self.sampler is not None: return len(self.sampler) // self.batch_size diff --git a/norch/utils/data/distributed.py b/norch/utils/data/distributed.py index c83dfc4..5d9f294 100644 --- a/norch/utils/data/distributed.py +++ b/norch/utils/data/distributed.py @@ -14,6 +14,8 @@ class DistributedSampler: # 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 diff --git a/tests/test_distributed.py b/tests/test_distributed.py index 86a429d..342d87f 100644 --- a/tests/test_distributed.py +++ b/tests/test_distributed.py @@ -4,11 +4,12 @@ from norch.utils import utils_unittests as utils import torch import torchvision import numpy as np +import os from norch.norchvision import transforms as norch_transforms from torchvision import transforms as torch_transforms class TestDistributed(unittest.TestCase): - def test_distributed_sampler_batch_1(self): + """def test_distributed_sampler_batch_1(self): transforms = norch_transforms.Compose( [ norch_transforms.ToTensor(), @@ -16,7 +17,7 @@ class TestDistributed(unittest.TestCase): ) train_data, test_data = norch.norchvision.datasets.MNIST.splits(transform=transforms, target_transform=transforms) distributed_sampler = norch.utils.data.distributed.DistributedSampler(dataset=train_data, num_replicas=8, rank=2) - train_loader = norch.utils.data.DataLoader(train_data, batch_size = 1, sampler=distributed_sampler) + train_loader = norch.utils.data.DataLoader(train_data, batch_size=1, sampler=distributed_sampler) labels_norch = [] for i, batch in enumerate(train_loader): @@ -24,7 +25,7 @@ class TestDistributed(unittest.TestCase): image, label = batch labels_norch.append(utils.to_torch(label)) - if i > 10: + if i > 100: break transforms = torch_transforms.Compose( @@ -35,7 +36,7 @@ class TestDistributed(unittest.TestCase): train_data = torchvision.datasets.MNIST(root='./.data/', download=True, transform=transforms) distributed_sampler = torch.utils.data.distributed.DistributedSampler(dataset=train_data, num_replicas=8, rank=2, shuffle=False) - train_loader = torch.utils.data.DataLoader(train_data, batch_size = 1, sampler=distributed_sampler) + train_loader = torch.utils.data.DataLoader(train_data, batch_size=1, sampler=distributed_sampler) labels_torch = [] for i, batch in enumerate(train_loader): @@ -43,22 +44,63 @@ class TestDistributed(unittest.TestCase): image, label = batch labels_torch.append(label) - if i > 10: + if i > 100: + break + + for label_norch, label_torch in zip(labels_norch, labels_torch): + self.assertTrue(utils.compare_torch(label_norch, label_torch))""" + + + def test_distributed_sampler_batch_32_rank_2(self): + transforms = norch_transforms.Compose( + [ + norch_transforms.ToTensor(), + ] + ) + train_data, test_data = norch.norchvision.datasets.MNIST.splits(transform=transforms, target_transform=transforms) + distributed_sampler = norch.utils.data.distributed.DistributedSampler(dataset=train_data, num_replicas=8, rank=2) + train_loader = norch.utils.data.DataLoader(train_data, batch_size=32, sampler=distributed_sampler) + + labels_norch = [] + for i, batch in enumerate(train_loader): + + image, label = batch + labels_norch.append(utils.to_torch(label)) + + if i > 100: + break + + transforms = torch_transforms.Compose( + [ + torch_transforms.ToTensor(), + ] + ) + + train_data = torchvision.datasets.MNIST(root='./.data/', download=True, transform=transforms) + distributed_sampler = torch.utils.data.distributed.DistributedSampler(dataset=train_data, num_replicas=8, rank=2, shuffle=False) + train_loader = torch.utils.data.DataLoader(train_data, batch_size=32, sampler=distributed_sampler) + + labels_torch = [] + for i, batch in enumerate(train_loader): + + image, label = batch + labels_torch.append(label) + + if i > 100: break for label_norch, label_torch in zip(labels_norch, labels_torch): self.assertTrue(utils.compare_torch(label_norch, label_torch)) - - def test_distributed_sampler_batch_32(self): + def test_distributed_sampler_batch_32_rank_7(self): transforms = norch_transforms.Compose( [ norch_transforms.ToTensor(), ] ) train_data, test_data = norch.norchvision.datasets.MNIST.splits(transform=transforms, target_transform=transforms) - distributed_sampler = norch.utils.data.distributed.DistributedSampler(dataset=train_data, num_replicas=8, rank=2) - train_loader = norch.utils.data.DataLoader(train_data, batch_size = 32, sampler=distributed_sampler) + distributed_sampler = norch.utils.data.distributed.DistributedSampler(dataset=train_data, num_replicas=8, rank=7) + train_loader = norch.utils.data.DataLoader(train_data, batch_size=32, sampler=distributed_sampler) labels_norch = [] for i, batch in enumerate(train_loader): @@ -66,7 +108,7 @@ class TestDistributed(unittest.TestCase): image, label = batch labels_norch.append(utils.to_torch(label)) - if i > 10: + if i > 100: break transforms = torch_transforms.Compose( @@ -76,20 +118,19 @@ class TestDistributed(unittest.TestCase): ) train_data = torchvision.datasets.MNIST(root='./.data/', download=True, transform=transforms) - distributed_sampler = torch.utils.data.distributed.DistributedSampler(dataset=train_data, num_replicas=8, rank=2, shuffle=False) - train_loader = torch.utils.data.DataLoader(train_data, batch_size = 32, sampler=distributed_sampler) - + distributed_sampler = torch.utils.data.distributed.DistributedSampler(dataset=train_data, num_replicas=8, rank=7, shuffle=False) + train_loader = torch.utils.data.DataLoader(train_data, batch_size=32, sampler=distributed_sampler) + labels_torch = [] for i, batch in enumerate(train_loader): image, label = batch labels_torch.append(label) - if i > 10: + if i > 100: break for label_norch, label_torch in zip(labels_norch, labels_torch): - print(label_norch, label_torch) self.assertTrue(utils.compare_torch(label_norch, label_torch)) diff --git a/train.py b/train.py index e999667..9ad1187 100644 --- a/train.py +++ b/train.py @@ -39,7 +39,7 @@ from norch.norchvision import transforms train_data, test_data = norch.norchvision.datasets.MNIST.splits(transform=transforms.ToTensor()) train_sampler = norch.utils.data.distributed.DistributedSampler(dataset=train_data, num_replicas=10, rank=2) -train_loader = norch.utils.data.Dataloader(train_data, batch_size = 1, sampler=train_sampler) +train_loader = norch.utils.data.Dataloader(train_data, batch_size = 50, sampler=train_sampler) input_sample, target_sample = train_data[0] fig = plt.figure(figsize = (20, 10))