文章目錄
- 一、影像增廣
- 二、常用的影像增廣方法
- 1. 翻轉和裁減
- 2. 顏色改變
- 3. 疊加使用多種資料增廣方法
- 三、使用影像增廣進行訓練
- 四、總結
一、影像增廣
定義&解釋:
- 通過對訓練影像做一系列隨機改變,來產生相似但又不同的訓練樣本,從而擴大訓練資料集的規模,
- 隨機改變訓練樣本可以降低模型對某些屬性的依賴,從而提高模型的范化能力
二、常用的影像增廣方法
使用下面這張400x500的影像作為范例
%matplotlib inline
import matplotlib.pyplot as plt
import numpy as np
import torch
import torchvision
from d2l import torch as d2l
from torch import nn
from PIL import Image
img = Image.open('./data/cat_dog/cat1.jpg')
plt.figure("cat")
plt.title('Initial data')
plt.imshow(img)
plt.show()
大多數影像增廣方法都具有一定的隨機性,為了便于觀察影像增廣的效果,我們下面定義輔助函式 apply , 此函式在輸入影像 img 上多次運行影像增廣方法 aug 并顯示所有結果,
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)
1. 翻轉和裁減
左右翻轉影像通常不會改變物件的類別,這是最早和最廣泛使用的影像增廣方法之一,上下翻轉影像不如左右影像翻轉那樣常用,但是,至少對于這個示例影像,上下翻轉不會妨礙識別,隨機裁減]在我們使用的示例影像中,貓位于影像的中間,但并非所有影像都是這樣, 池化層可以降低卷積層對目標位置的敏感性, 另外,我們可以通過對影像進行隨機裁剪,使物體以不同的比例出現在影像的不同位置, 這也可以降低模型對目標位置的敏感性,
# 左右翻轉
apply(img,torchvision.transforms.RandomHorizontalFlip())
# 上下翻轉
apply(img,torchvision.transforms.RandomVerticalFlip())
# 隨機裁減
shape_aug = torchvision.transforms.RandomResizedCrop(
(200,200),scale=(0.1,1),ratio=(0.5,2),
# (200,200)是圖片的大小,scale表示隨機裁減為原來的比例,ratio是長寬比
)
apply(img,shape_aug)
2. 顏色改變
另一種增廣方法是改變顏色,
我們可以改變影像顏色的四個方面:
- 亮度
- 對比度
- 飽和度
- 色調
# 亮度
apply(img,
torchvision.transforms.ColorJitter(brightness=0.5,contrast=0,
saturation=0,hue=0))
# 對比度
apply(img,
torchvision.transforms.ColorJitter(brightness=0,contrast=0.5,
saturation=0,hue=0))
# 飽和度
apply(img,
torchvision.transforms.ColorJitter(brightness=0,contrast=0,
saturation=0.5,hue=0))
# 色調
apply(img,
torchvision.transforms.ColorJitter(brightness=0,contrast=0,
saturation=0,hue=0.5))
# 混合使用
apply(img,
torchvision.transforms.ColorJitter(brightness=0.5,contrast=0.5,
saturation=0.5,hue=0.5))
3. 疊加使用多種資料增廣方法
augs = torchvision.transforms.Compose(
[torchvision.transforms.RandomHorizontalFlip(),
torchvision.transforms.ColorJitter(brightness=0.5,contrast=0.5,saturation=0.5,hue=0.5),
torchvision.transforms.RandomResizedCrop((200, 200), scale=(0.1, 1), ratio=(0.5, 2))
]
)
apply(img,augs)
三、使用影像增廣進行訓練
# 下載CIFA10資料集測驗
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)
# 應用簡單的左右翻轉,上下翻轉
# 生資料格式為(批量大小,通道數量,高度,寬度)
train_augs = torchvision.transforms.Compose(
[torchvision.transforms.RandomHorizontalFlip(),
# torchvision.transforms.RandomVerticalFlip(),
torchvision.transforms.ToTensor()]
)
test_augs = torchvision.transforms.Compose([
# torchvision.transforms.RandomHorizontalFlip(),
# torchvision.transforms.RandomVerticalFlip(),
torchvision.transforms.ToTensor()
]
)
# 加載資料
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 dataLoader
# 多GPU訓練和評估
def train_batch(net, X, y, loss, trainer, devices):
if isinstance(X, list):
# 微調BERT中所需(稍后討論)
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_sum
def train(net, train_iter, test_iter, loss, trainer, num_epochs,
devices=d2l.try_all_gpus()):
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]) # 多GPU運行
for epoch in range(num_epochs):
# 4個維度:儲存訓練損失,訓練準確度,實體數,特點數
metric = d2l.Accumulator(4)
for i, (features, labels) in enumerate(train_iter):
timer.start()
l, acc = train_batch(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)}')
#使用增強之后的資料進行訓練模型;
# 獲取全部的GPU,使用Adam作為優化演算法
batch_size, devices, net = 256, d2l.try_all_gpus(), d2l.resnet18(10, 3)
# 模型初始化
def init_weights(m):
if type(m) in [nn.Linear, nn.Conv2d]:
nn.init.xavier_uniform_(m.weight)
net.apply(init_weights)
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)
train(net, train_iter, test_iter, loss, trainer, 10, devices)
# 資料增廣(左右翻轉)
train_with_data_aug(train_augs, test_augs, net)
loss 0.166, train acc 0.942, test acc 0.823
453.4 examples/sec on [device(type='cuda', index=0)]
# 沒有資料增廣
batch_size, devices, net = 256, d2l.try_all_gpus(), d2l.resnet18(10, 3)
def init_weights(m):
if type(m) in [nn.Linear, nn.Conv2d]:
nn.init.xavier_uniform_(m.weight)
net.apply(init_weights)
train_with_data_aug(test_augs, test_augs, net)
loss 0.070, train acc 0.975, test acc 0.797
455.5 examples/sec on [device(type='cuda', index=0)]
結果對比:
- 使用影像增強,盡管只是簡單的左右翻轉,我們模型的預測精度還是提高了3%
- 模型過擬合有一定的緩解,
四、總結
- 影像增廣基于現有的訓練資料生成隨機影像,來提高模型的范化能力,
- 為了在預測程序中得到確切的結果,我們通常對訓練樣本只進行影像增廣,而在預測程序中不使用隨機操作的影像增廣,(訓練有,預測無)
- 深度學習框架提供了許多不同的影像增廣方法,這些方法可以被同時應用,(多種增強共同使用)
- 影像增廣方法收集(這些整理應該夠用了,如果有什么特別需求可以留言討論一下):
(1)知乎上有作者總結自己撰寫的15種增強方法和代碼:
https://zhuanlan.zhihu.com/p/158854758- 翻轉
- 裁剪
- 過濾和銳化
- 模糊
- 旋轉,平移,剪切,縮放
- 剪下
- 色彩
- 亮度
- 對比
- 均勻和高斯噪聲
- 漸變鏡頭變形
(2)github上找一些高star的成熟代碼:
例如: imgaug https://github.com/aleju/imgaug
(3)augmentor https://github.com/mdbloice/Augmentor
轉載請註明出處,本文鏈接:https://www.uj5u.com/qita/296685.html
標籤:其他
上一篇:深度學習之影像分類(五)--GoogLeNet網路結構
下一篇:基于MATLAB的空心散點檢測
