1.8 MiB
1.8 MiB
In [1]:
%matplotlib inline
import torch
import torchvision
from torch import nn
from d2l import torch as d2lIn [2]:
d2l.set_figsize()
img = d2l.Image.open('../img/cat1.jpg')
d2l.plt.imshow(img);In [3]:
def apply(img, aug, num_rows=2, num_cols=4, scale=1.5):
Y = [aug(img) for _ in range(num_rows * num_cols)]
d2l.show_images(Y, num_rows, num_cols, scale=scale)In [4]:
apply(img, torchvision.transforms.RandomHorizontalFlip())In [5]:
apply(img, torchvision.transforms.RandomVerticalFlip())In [6]:
shape_aug = torchvision.transforms.RandomResizedCrop(
(200, 200), scale=(0.1, 1), ratio=(0.5, 2))
apply(img, shape_aug)In [7]:
apply(img, torchvision.transforms.ColorJitter(
brightness=0.5, contrast=0, saturation=0, hue=0))In [8]:
apply(img, torchvision.transforms.ColorJitter(
brightness=0, contrast=0, saturation=0, hue=0.5))In [9]:
color_aug = torchvision.transforms.ColorJitter(
brightness=0.5, contrast=0.5, saturation=0.5, hue=0.5)
apply(img, color_aug)In [10]:
augs = torchvision.transforms.Compose([
torchvision.transforms.RandomHorizontalFlip(), color_aug, shape_aug])
apply(img, augs)In [11]:
all_images = torchvision.datasets.CIFAR10(train=True, root="../data",
download=True)
d2l.show_images([all_images[i][0] for i in range(32)], 4, 8, scale=0.8);Downloading https://www.cs.toronto.edu/~kriz/cifar-10-python.tar.gz to ../data/cifar-10-python.tar.gz
0%| | 0/170498071 [00:00<?, ?it/s]
0%| | 65536/170498071 [00:00<06:06, 465602.77it/s]
0%| | 229376/170498071 [00:00<03:15, 870075.02it/s]
1%| | 917504/170498071 [00:00<01:03, 2680566.44it/s]
2%|▏ | 3604480/170498071 [00:00<00:18, 9055226.73it/s]
6%|▌ | 9601024/170498071 [00:00<00:06, 23469988.96it/s]
7%|▋ | 12779520/170498071 [00:00<00:06, 24635235.07it/s]
11%|█ | 18087936/170498071 [00:00<00:04, 32644362.62it/s]
13%|█▎ | 22085632/170498071 [00:00<00:04, 33766951.34it/s]
15%|█▌ | 25919488/170498071 [00:01<00:04, 34874971.42it/s]
18%|█▊ | 31424512/170498071 [00:01<00:03, 39206606.36it/s]
21%|██ | 35454976/170498071 [00:01<00:03, 38026146.88it/s]
24%|██▍ | 40730624/170498071 [00:01<00:03, 42003142.31it/s]
26%|██▋ | 45023232/170498071 [00:01<00:03, 39377573.03it/s]
29%|██▉ | 50266112/170498071 [00:01<00:02, 42853627.52it/s]
32%|███▏ | 54657024/170498071 [00:01<00:02, 40134646.49it/s]
35%|███▌ | 59998208/170498071 [00:01<00:02, 43645137.55it/s]
38%|███▊ | 64487424/170498071 [00:02<00:02, 40692269.19it/s]
41%|████ | 69599232/170498071 [00:02<00:02, 43277945.93it/s]
43%|████▎ | 74022912/170498071 [00:02<00:02, 40757798.12it/s]
46%|████▋ | 78970880/170498071 [00:02<00:02, 42534999.51it/s]
49%|████▉ | 83296256/170498071 [00:02<00:02, 40565223.57it/s]
52%|█████▏ | 88408064/170498071 [00:02<00:01, 43070011.08it/s]
54%|█████▍ | 92798976/170498071 [00:02<00:01, 40641250.05it/s]
57%|█████▋ | 97943552/170498071 [00:02<00:01, 43576448.82it/s]
60%|██████ | 102400000/170498071 [00:02<00:01, 40609636.82it/s]
63%|██████▎ | 107577344/170498071 [00:03<00:01, 43606965.98it/s]
66%|██████▌ | 112033792/170498071 [00:03<00:01, 40736940.01it/s]
69%|██████▊ | 117112832/170498071 [00:03<00:01, 43233593.74it/s]
71%|███████▏ | 121536512/170498071 [00:03<00:01, 40710642.41it/s]
74%|███████▍ | 126615552/170498071 [00:03<00:01, 43290286.82it/s]
77%|███████▋ | 131039232/170498071 [00:03<00:00, 40679159.83it/s]
80%|███████▉ | 136118272/170498071 [00:03<00:00, 43195601.76it/s]
82%|████████▏ | 140541952/170498071 [00:03<00:00, 40621722.11it/s]
85%|████████▌ | 145555456/170498071 [00:03<00:00, 43050443.46it/s]
88%|████████▊ | 149946368/170498071 [00:04<00:00, 40446858.58it/s]
91%|█████████ | 155123712/170498071 [00:04<00:00, 43392878.70it/s]
94%|█████████▎| 159547392/170498071 [00:04<00:00, 40552440.46it/s]
97%|█████████▋| 164823040/170498071 [00:04<00:00, 43601133.14it/s]
99%|█████████▉| 169279488/170498071 [00:04<00:00, 40752923.51it/s]
100%|██████████| 170498071/170498071 [00:04<00:00, 37716809.52it/s]
Extracting ../data/cifar-10-python.tar.gz to ../data
In [12]:
train_augs = torchvision.transforms.Compose([
torchvision.transforms.RandomHorizontalFlip(),
torchvision.transforms.ToTensor()])
test_augs = torchvision.transforms.Compose([
torchvision.transforms.ToTensor()])In [13]:
def load_cifar10(is_train, augs, batch_size):
dataset = torchvision.datasets.CIFAR10(root="../data", train=is_train,
transform=augs, download=True)
dataloader = torch.utils.data.DataLoader(dataset, batch_size=batch_size,
shuffle=is_train, num_workers=d2l.get_dataloader_workers())
return dataloaderIn [14]:
#@save
def train_batch_ch13(net, X, y, loss, trainer, devices):
"""Train for a minibatch with multiple GPUs (defined in Chapter 13)."""
if isinstance(X, list):
# Required for BERT fine-tuning (to be covered later)
X = [x.to(devices[0]) for x in X]
else:
X = X.to(devices[0])
y = y.to(devices[0])
net.train()
trainer.zero_grad()
pred = net(X)
l = loss(pred, y)
l.sum().backward()
trainer.step()
train_loss_sum = l.sum()
train_acc_sum = d2l.accuracy(pred, y)
return train_loss_sum, train_acc_sumIn [15]:
#@save
def train_ch13(net, train_iter, test_iter, loss, trainer, num_epochs,
devices=d2l.try_all_gpus()):
"""Train a model with multiple GPUs (defined in Chapter 13)."""
timer, num_batches = d2l.Timer(), len(train_iter)
animator = d2l.Animator(xlabel='epoch', xlim=[1, num_epochs], ylim=[0, 1],
legend=['train loss', 'train acc', 'test acc'])
net = nn.DataParallel(net, device_ids=devices).to(devices[0])
for epoch in range(num_epochs):
# Sum of training loss, sum of training accuracy, no. of examples,
# no. of predictions
metric = d2l.Accumulator(4)
for i, (features, labels) in enumerate(train_iter):
timer.start()
l, acc = train_batch_ch13(
net, features, labels, loss, trainer, devices)
metric.add(l, acc, labels.shape[0], labels.numel())
timer.stop()
if (i + 1) % (num_batches // 5) == 0 or i == num_batches - 1:
animator.add(epoch + (i + 1) / num_batches,
(metric[0] / metric[2], metric[1] / metric[3],
None))
test_acc = d2l.evaluate_accuracy_gpu(net, test_iter)
animator.add(epoch + 1, (None, None, test_acc))
print(f'loss {metric[0] / metric[2]:.3f}, train acc '
f'{metric[1] / metric[3]:.3f}, test acc {test_acc:.3f}')
print(f'{metric[2] * num_epochs / timer.sum():.1f} examples/sec on '
f'{str(devices)}')In [16]:
batch_size, devices, net = 256, d2l.try_all_gpus(), d2l.resnet18(10, 3)
net.apply(d2l.init_cnn)
def train_with_data_aug(train_augs, test_augs, net, lr=0.001):
train_iter = load_cifar10(True, train_augs, batch_size)
test_iter = load_cifar10(False, test_augs, batch_size)
loss = nn.CrossEntropyLoss(reduction="none")
trainer = torch.optim.Adam(net.parameters(), lr=lr)
net(next(iter(train_iter))[0])
train_ch13(net, train_iter, test_iter, loss, trainer, 10, devices)In [17]:
train_with_data_aug(train_augs, test_augs, net)loss 0.215, train acc 0.925, test acc 0.810 4728.8 examples/sec on [device(type='cuda', index=0), device(type='cuda', index=1)]