fix distributed dataloader
This commit is contained in:
parent
10a88fe8ca
commit
e3bb882d3c
4 changed files with 69 additions and 27 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
||||
|
||||
|
|
|
|||
2
train.py
2
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))
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue