commit
f5d0a1271d
25 changed files with 381 additions and 35 deletions
10
README.md
10
README.md
|
|
@ -65,7 +65,7 @@ class MyModel(nn.Module):
|
|||
### 3.3 - Example training
|
||||
```python
|
||||
import norch
|
||||
from norch.utils.data.dataloader import Dataloader
|
||||
from norch.utils.data.dataloader import DataLoader
|
||||
from norch.norchvision import transforms
|
||||
import norch
|
||||
import norch.nn as nn
|
||||
|
|
@ -77,21 +77,21 @@ BATCH_SIZE = 32
|
|||
device = "cuda" #cpu
|
||||
epochs = 10
|
||||
|
||||
transform = transforms.Sequential(
|
||||
transform = transforms.Compose(
|
||||
[
|
||||
transforms.ToTensor(),
|
||||
transforms.Reshape([-1, 784, 1])
|
||||
]
|
||||
)
|
||||
|
||||
target_transform = transforms.Sequential(
|
||||
target_transform = transforms.Compose(
|
||||
[
|
||||
transforms.ToTensor()
|
||||
]
|
||||
)
|
||||
|
||||
train_data, test_data = norch.norchvision.datasets.MNIST.splits(transform=transform, target_transform=target_transform)
|
||||
train_loader = Dataloader(train_data, batch_size = BATCH_SIZE)
|
||||
train_loader = DataLoader(train_data, batch_size = BATCH_SIZE)
|
||||
|
||||
class MyModel(nn.Module):
|
||||
def __init__(self):
|
||||
|
|
@ -146,4 +146,4 @@ for epoch in range(epochs):
|
|||
| Loss | in progress | <ul><li>[x] MSE</li><li>[X] Cross Entropy</li></ul> |
|
||||
| Data | in progress | <ul><li>[X] Dataset</li><li>[X] Batch</li><li>[X] Iterator</li></ul> |
|
||||
| Convolutional Neural Network | in progress | <ul><li>[ ] Conv2d</li><li>[ ] MaxPool2d</li><li>[ ] Dropout</li></ul> |
|
||||
| Distributed | in progress | <ul><li>[ ] Distributed Data Parallel</li></ul>
|
||||
| Distributed | in progress | <ul>><li>[ ] All reduce</li><li>[ ] DistributedDataParallel</li>><li>[ ] DistributedSampler</li></ul>
|
||||
|
|
|
|||
Binary file not shown.
|
|
@ -33,7 +33,7 @@
|
|||
"import norch\n",
|
||||
"import norch.nn as nn\n",
|
||||
"import norch.optim as optim\n",
|
||||
"from norch.utils.data.dataloader import Dataloader\n",
|
||||
"from norch.utils.data.dataloader import DataLoader\n",
|
||||
"from norch.norchvision import transforms\n",
|
||||
"import numpy as np\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
|
|
@ -102,7 +102,7 @@
|
|||
"device = \"cuda\" #cpu\n",
|
||||
"epochs = 10\n",
|
||||
"\n",
|
||||
"transform = transforms.Sequential(\n",
|
||||
"transform = transforms.Compose(\n",
|
||||
" [\n",
|
||||
" transforms.ToTensor(),\n",
|
||||
" transforms.Reshape([-1, 784, 1])\n",
|
||||
|
|
@ -116,7 +116,7 @@
|
|||
")\n",
|
||||
"\n",
|
||||
"train_data, test_data = norch.norchvision.datasets.MNIST.splits(transform=transform, target_transform=target_transform)\n",
|
||||
"train_loader = Dataloader(train_data, batch_size = BATCH_SIZE)\n",
|
||||
"train_loader = DataLoader(train_data, batch_size = BATCH_SIZE)\n",
|
||||
"\n",
|
||||
"class MyModel(nn.Module):\n",
|
||||
" def __init__(self):\n",
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ from .nn import *
|
|||
from .optim import *
|
||||
from .utils import *
|
||||
from .norchvision import *
|
||||
from .utils import *
|
||||
|
||||
__version__ = "0.0.4"
|
||||
__author__ = 'Lucas de Lima Nogueira'
|
||||
|
|
|
|||
Binary file not shown.
|
|
@ -35,12 +35,12 @@ void init_process_group(int env_rank, int env_world_size) {
|
|||
|
||||
}
|
||||
|
||||
void broadcast_tensor(Tensor* tensor) {
|
||||
void broadcast_tensor(Tensor* tensor, int src) {
|
||||
|
||||
cudaStream_t stream;
|
||||
cudaStreamCreate(&stream);
|
||||
|
||||
NCCL_CHECK(ncclBroadcast(tensor->data, tensor->data, tensor->size * sizeof(float), ncclFloat, 0, nccl_comm, stream));
|
||||
NCCL_CHECK(ncclBroadcast(tensor->data, tensor->data, tensor->size * sizeof(float), ncclFloat, src, nccl_comm, stream));
|
||||
cudaStreamSynchronize(stream);
|
||||
cudaStreamDestroy(stream);
|
||||
}
|
||||
|
|
@ -56,6 +56,19 @@ void allreduce_sum_tensor(Tensor* tensor) {
|
|||
cudaStreamDestroy(stream);
|
||||
}
|
||||
|
||||
void allreduce_mean_tensor(Tensor* tensor) {
|
||||
cudaStream_t stream;
|
||||
cudaStreamCreate(&stream);
|
||||
|
||||
// Perform NCCL AllReduce operation to calculate the mean of all tensors across all processes
|
||||
NCCL_CHECK(ncclAllReduce(tensor->data, tensor->data, tensor->size, ncclFloat, ncclSum, nccl_comm, stream));
|
||||
|
||||
tensor_div_scalar_cuda(tensor, world_size, tensor->data);
|
||||
|
||||
cudaStreamSynchronize(stream);
|
||||
cudaStreamDestroy(stream);
|
||||
}
|
||||
|
||||
void end_process_group() {
|
||||
MPI_CHECK(MPI_Finalize());
|
||||
NCCL_CHECK(ncclCommDestroy(nccl_comm));
|
||||
|
|
|
|||
|
|
@ -22,8 +22,9 @@
|
|||
|
||||
extern "C" {
|
||||
void init_process_group(int rank, int world_size);
|
||||
void broadcast_tensor(Tensor* tensor);
|
||||
void broadcast_tensor(Tensor* tensor, int src);
|
||||
void allreduce_sum_tensor(Tensor* tensor);
|
||||
void allreduce_mean_tensor(Tensor* tensor);
|
||||
}
|
||||
|
||||
#endif /* DISTRIBUTED_H */
|
||||
|
|
|
|||
|
|
@ -9,12 +9,18 @@ def init_process_group(rank, world_size, backend='nccl'):
|
|||
|
||||
Tensor._C.init_process_group(rank, world_size)
|
||||
|
||||
def broadcast_tensor(tensor):
|
||||
def get_rank():
|
||||
return os.getenv('OMPI_COMM_WORLD_RANK', 0)
|
||||
|
||||
Tensor._C.broadcast_tensor.argtypes = [ctypes.POINTER(CTensor)]
|
||||
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]
|
||||
Tensor._C.broadcast_tensor.restype = None
|
||||
|
||||
Tensor._C.broadcast_tensor(tensor.tensor)
|
||||
Tensor._C.broadcast_tensor(tensor.tensor, src)
|
||||
|
||||
def allreduce_sum_tensor(tensor):
|
||||
|
||||
|
|
@ -22,3 +28,10 @@ def allreduce_sum_tensor(tensor):
|
|||
Tensor._C.allreduce_sum_tensor.restype = None
|
||||
|
||||
Tensor._C.allreduce_sum_tensor(tensor.tensor)
|
||||
|
||||
def allreduce_mean_tensor(tensor):
|
||||
|
||||
Tensor._C.allreduce_mean_tensor.argtypes = [ctypes.POINTER(CTensor)]
|
||||
Tensor._C.allreduce_mean_tensor.restype = None
|
||||
|
||||
Tensor._C.allreduce_mean_tensor(tensor.tensor)
|
||||
|
|
|
|||
Binary file not shown.
|
|
@ -1,4 +1,5 @@
|
|||
from .modules import *
|
||||
from .activation import *
|
||||
from .loss import *
|
||||
from .functional import *
|
||||
from .functional import *
|
||||
from .parallel import *
|
||||
Binary file not shown.
Binary file not shown.
21
norch/nn/parallel.py
Normal file
21
norch/nn/parallel.py
Normal file
|
|
@ -0,0 +1,21 @@
|
|||
from .module import *
|
||||
import norch.distributed as dist
|
||||
import os
|
||||
|
||||
class DistributedDataParallel(Module):
|
||||
def __init__(self, module):
|
||||
super().__init__()
|
||||
|
||||
self.module = module
|
||||
|
||||
def forward(self, *inputs, **kwargs):
|
||||
return self.module(*inputs, **kwargs)
|
||||
|
||||
def backward(self):
|
||||
self.module.backward()
|
||||
for module, name, parameter in self.parameters():
|
||||
dist.allreduce_mean_tensor(parameter)
|
||||
|
||||
|
||||
|
||||
|
||||
|
|
@ -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)
|
||||
|
|
@ -11,7 +11,7 @@ class Reshape:
|
|||
def __call__(self, x):
|
||||
return x.reshape(self.shape)
|
||||
|
||||
class Sequential:
|
||||
class Compose:
|
||||
def __init__(self, transforms):
|
||||
self.transforms = transforms
|
||||
|
||||
|
|
|
|||
|
|
@ -1,2 +1,2 @@
|
|||
from .utils import *
|
||||
from .data import *
|
||||
from .data import *
|
||||
from .functions import *
|
||||
|
|
@ -1,4 +1,5 @@
|
|||
from .dataset import *
|
||||
from .example import *
|
||||
from .dataloader import *
|
||||
from .batch import *
|
||||
from .batch import *
|
||||
from .distributed import *
|
||||
|
|
@ -2,19 +2,27 @@ 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, 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:
|
||||
end = start + self.batch_size
|
||||
batch_size = min(end, len(self.dataset)) - start
|
||||
yield Batch(self.dataset[start:end], batch_size)
|
||||
|
||||
if self.sampler is not None:
|
||||
indices = list(iter(self.sampler))
|
||||
|
||||
else:
|
||||
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
|
||||
|
||||
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
|
||||
|
||||
|
|
|
|||
28
norch/utils/data/distributed.py
Normal file
28
norch/utils/data/distributed.py
Normal file
|
|
@ -0,0 +1,28 @@
|
|||
import numpy as np
|
||||
import math
|
||||
|
||||
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):
|
||||
indices = list(range(len(self.dataset)))
|
||||
|
||||
# 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
|
||||
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
|
||||
93
tests/test_dataset.py
Normal file
93
tests/test_dataset.py
Normal file
|
|
@ -0,0 +1,93 @@
|
|||
import unittest
|
||||
import norch
|
||||
from norch.utils import utils_unittests as utils
|
||||
import torch
|
||||
import torchvision
|
||||
import numpy as np
|
||||
from norch.norchvision import transforms as norch_transforms
|
||||
from torchvision import transforms as torch_transforms
|
||||
|
||||
class TestDataset(unittest.TestCase):
|
||||
def test_dataloader_batchsize_1(self):
|
||||
transforms = norch_transforms.Compose(
|
||||
[
|
||||
norch_transforms.ToTensor(),
|
||||
]
|
||||
)
|
||||
train_data, test_data = norch.norchvision.datasets.MNIST.splits(transform=transforms, target_transform=transforms)
|
||||
train_loader = norch.utils.data.DataLoader(train_data, batch_size = 1)
|
||||
|
||||
labels_norch = []
|
||||
for i, batch in enumerate(train_loader):
|
||||
|
||||
image, label = batch
|
||||
labels_norch.append(utils.to_torch(label))
|
||||
|
||||
if i > 10:
|
||||
break
|
||||
|
||||
transforms = torch_transforms.Compose(
|
||||
[
|
||||
torch_transforms.ToTensor(),
|
||||
]
|
||||
)
|
||||
|
||||
train_data = torchvision.datasets.MNIST(root='./.data/', download=True, transform=transforms)
|
||||
train_loader = torch.utils.data.DataLoader(train_data, batch_size = 1)
|
||||
|
||||
labels_torch = []
|
||||
for i, batch in enumerate(train_loader):
|
||||
|
||||
image, label = batch
|
||||
labels_torch.append(label)
|
||||
|
||||
if i > 10:
|
||||
break
|
||||
|
||||
for label_norch, label_torch in zip(labels_norch, labels_torch):
|
||||
self.assertTrue(utils.compare_torch(label_norch, label_torch))
|
||||
|
||||
|
||||
def test_dataloader_batchsize_32(self):
|
||||
transforms = norch_transforms.Compose(
|
||||
[
|
||||
norch_transforms.ToTensor(),
|
||||
]
|
||||
)
|
||||
train_data, test_data = norch.norchvision.datasets.MNIST.splits(transform=transforms, target_transform=transforms)
|
||||
train_loader = norch.utils.data.DataLoader(train_data, batch_size = 32)
|
||||
|
||||
labels_norch = []
|
||||
for i, batch in enumerate(train_loader):
|
||||
|
||||
image, label = batch
|
||||
labels_norch.append(utils.to_torch(label))
|
||||
|
||||
if i > 10:
|
||||
break
|
||||
|
||||
transforms = torch_transforms.Compose(
|
||||
[
|
||||
torch_transforms.ToTensor(),
|
||||
]
|
||||
)
|
||||
|
||||
train_data = torchvision.datasets.MNIST(root='./.data/', download=True, transform=transforms)
|
||||
train_loader = torch.utils.data.DataLoader(train_data, batch_size = 32)
|
||||
|
||||
labels_torch = []
|
||||
for i, batch in enumerate(train_loader):
|
||||
|
||||
image, label = batch
|
||||
labels_torch.append(label)
|
||||
|
||||
if i > 10:
|
||||
break
|
||||
|
||||
for label_norch, label_torch in zip(labels_norch, labels_torch):
|
||||
self.assertTrue(utils.compare_torch(label_norch, label_torch))
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
139
tests/test_distributed.py
Normal file
139
tests/test_distributed.py
Normal file
|
|
@ -0,0 +1,139 @@
|
|||
import unittest
|
||||
import norch
|
||||
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):
|
||||
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=1, 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=1, 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_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_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=7)
|
||||
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=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 > 100:
|
||||
break
|
||||
|
||||
for label_norch, label_torch in zip(labels_norch, labels_torch):
|
||||
self.assertTrue(utils.compare_torch(label_norch, label_torch))
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
|
@ -200,5 +200,4 @@ class TestNNModuleActivationFn(unittest.TestCase):
|
|||
softmax_torch_expected = softmax_fn_torch.forward(torch_input)
|
||||
|
||||
# Compare the results
|
||||
self.assertTrue(utils.compare_torch(softmax_torch_result, softmax_torch_expected))
|
||||
|
||||
self.assertTrue(utils.compare_torch(softmax_torch_result, softmax_torch_expected))
|
||||
30
train.py
30
train.py
|
|
@ -1,4 +1,4 @@
|
|||
import os
|
||||
"""import os
|
||||
import norch
|
||||
import norch.distributed as dist
|
||||
|
||||
|
|
@ -29,3 +29,31 @@ def main():
|
|||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
"""
|
||||
import norch
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
import random
|
||||
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 = 50, 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