implemented distributed sampler
This commit is contained in:
parent
56fff47001
commit
e590f57555
12 changed files with 56 additions and 22 deletions
|
|
@ -3,7 +3,7 @@ from .nn import *
|
|||
from .optim import *
|
||||
from .utils import *
|
||||
from .norchvision import *
|
||||
from . import data
|
||||
from .utils import *
|
||||
|
||||
__version__ = "0.0.4"
|
||||
__author__ = 'Lucas de Lima Nogueira'
|
||||
|
|
|
|||
Binary file not shown.
|
|
@ -9,6 +9,12 @@ def init_process_group(rank, world_size, backend='nccl'):
|
|||
|
||||
Tensor._C.init_process_group(rank, world_size)
|
||||
|
||||
def get_rank():
|
||||
return os.getenv('OMPI_COMM_WORLD_RANK', 0)
|
||||
|
||||
def get_world_size():
|
||||
return os.getenv('OMPI_COMM_WORLD_SIZE', 1)
|
||||
|
||||
def broadcast_tensor(tensor, src=0):
|
||||
|
||||
Tensor._C.broadcast_tensor.argtypes = [ctypes.POINTER(CTensor), ctypes.c_int]
|
||||
|
|
|
|||
Binary file not shown.
Binary file not shown.
|
|
@ -1,5 +1,5 @@
|
|||
from norch.tensor import Tensor
|
||||
from norch.utils import utils
|
||||
from norch.utils import functions
|
||||
import random
|
||||
|
||||
class Parameter(Tensor):
|
||||
|
|
@ -7,5 +7,5 @@ class Parameter(Tensor):
|
|||
A parameter is a trainable tensor.
|
||||
"""
|
||||
def __init__(self, shape):
|
||||
data = utils.generate_random_list(shape=shape)
|
||||
data = functions.generate_random_list(shape=shape)
|
||||
super().__init__(data, requires_grad=True)
|
||||
|
|
@ -1,2 +1,2 @@
|
|||
from .utils import *
|
||||
from .data import *
|
||||
from .data import *
|
||||
from .functions import *
|
||||
|
|
@ -11,17 +11,19 @@ class Dataloader:
|
|||
|
||||
def __iter__(self):
|
||||
if self.sampler is not None:
|
||||
indices = list(self.sampler)
|
||||
indices = iter(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_indices = indices[start:end]
|
||||
batch_size = len(batch_indices)
|
||||
yield Batch([self.dataset[i] for i in batch_indices], batch_size)
|
||||
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)
|
||||
|
||||
|
||||
def __len__(self):
|
||||
if self.sampler is not None:
|
||||
return len(self.sampler) // self.batch_size
|
||||
|
||||
return len(self.dataset) // self.batch_size
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
from abc import ABC, abstractmethod
|
||||
import os
|
||||
from norch.utils import extract_to_dir, download_from_url
|
||||
from norch.utils.functions import extract_to_dir, download_from_url
|
||||
from .example import Example
|
||||
import norch
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
import math
|
||||
import numpy as np
|
||||
import math
|
||||
|
||||
class DistributedSampler:
|
||||
def __init__(self, dataset, num_replicas, rank):
|
||||
|
|
@ -10,18 +10,17 @@ class DistributedSampler:
|
|||
self.total_size = self.num_samples * self.num_replicas
|
||||
|
||||
def __iter__(self):
|
||||
# Create indices for the dataset
|
||||
indices = np.arange(len(self.dataset))
|
||||
indices = list(range(len(self.dataset)))
|
||||
|
||||
# Add extra samples to make it evenly divisible
|
||||
indices = np.concatenate([indices, indices[:(self.total_size - len(indices)) % len(indices)]])
|
||||
indices = indices[:self.total_size]
|
||||
assert len(indices) == self.total_size
|
||||
|
||||
# Subsample
|
||||
# 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
|
||||
return self.num_samples
|
||||
|
|
|
|||
29
train.py
29
train.py
|
|
@ -1,4 +1,4 @@
|
|||
import os
|
||||
"""import os
|
||||
import norch
|
||||
import norch.distributed as dist
|
||||
|
||||
|
|
@ -29,3 +29,30 @@ def main():
|
|||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
"""
|
||||
import norch
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
import random
|
||||
|
||||
train_data, test_data = norch.norchvision.datasets.MNIST.splits()
|
||||
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)
|
||||
input_sample, target_sample = train_data[0]
|
||||
|
||||
fig = plt.figure(figsize = (20, 10))
|
||||
columns = 4
|
||||
rows = 2
|
||||
|
||||
# Choose a random image
|
||||
for image_index, batch in enumerate(train_loader):
|
||||
|
||||
image, label = batch
|
||||
fig.add_subplot(rows, columns, image_index+1)
|
||||
plt.imshow(np.array(image).reshape(28, 28))
|
||||
plt.title(label)
|
||||
plt.axis('off')
|
||||
if image_index > 6:
|
||||
break
|
||||
plt.show()
|
||||
Loading…
Add table
Add a link
Reference in a new issue