fix distributed dataloader

This commit is contained in:
lucasdelimanogueira 2024-05-27 15:08:27 -03:00
parent 10a88fe8ca
commit e3bb882d3c
4 changed files with 69 additions and 27 deletions

View file

@ -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

View file

@ -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

View file

@ -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))

View file

@ -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))