【Pytorch教程】Pytorch tutorials 04-Training a classifer 中文翻译

Training a classifier

本篇文章是本人对Pytorch官方教程的原创翻译(原文链接)仅供学习交流使用,转载请注明出处!

现在我们已经掌握了如何去定义神经网络、计算误差、更新权重。但在前面的章节中,我们用到的数据集都是自己构造的虚拟数据,那么如何真正地处理数据呢?

通常,我们处理图像、文本、音频、视频等数据时,可以使用一些Python的标准库,将输入导入为numpy格式,然后我们将导入的numpy数组转化为tensor。

  • 处理图像数据,用PillowOpenCV
  • 处理音频,用scipylibrosa
  • 处理文本,既可以使用Python/Cython的原生方法,也可以使用NLTKSpacy

pytorch为计算机视觉任务特别提供了一个torchvision包,内含Imagenet、CIFAR10、MNIST等常用数据集,以及数据集的转换器。他们分别包含在torchvision.datasetstorch.utils.data.DataLoader中。这就极大地避免了编写大量重复的代码。

本篇教程会使用CIFAR10数据集。它由10类图片组成,每张图片都是32x32,3通道像素。

Training an image classifier

创建一个图像分类器共需5个步骤:

  1. torchvision加载CIFAR10数据集并标准化。
  2. 定义一个卷积神经网络
  3. 定义损失函数
  4. 用训练集训练网络
  5. 用测试集测试网络

步骤1 加载CIFAR10数据集并标准化。

import torch
import torchvision
import torchvision.transforms as transforms

torchvision.datasets提供的图像是PILImage,像素在[0, 1] 区间,我们需要将其标准化,得到的是[-1, 1]的数据。

transform = transforms.Compose([transforms.ToTensor(),
                                 transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))])  # (样本-均值) / 标准差, 需要分别指定3个通道的均值和标注差

'''
加载训练集
root:数据集根目录
train:是否为训练集
download:是否需要下载
transform:transform对象,对数据集进行转换
'''
trainset = torchvision.datasets.CIFAR10(root='./data', train=True, download=True, transform=transform)
# shuffle:是否打乱 num_workers: 多线程数量 如果在windows下报错请改为0
trainloader = torch.utils.data.DataLoader(trainset, batch_size=4, shuffle=True, num_workers=2)

# 加载测试集,与上面同理
testset = torchvision.datasets.CIFAR10(root='./data', train=False, download=True, transform=transform)
testloader = torch.utils.data.DataLoader(testset, batch_size=4, shuffle=False, num_workers=2)

classes = ('plane', 'car', 'bird', 'cat','deer', 'dog', 'frog', 'horse', 'ship', 'truck')
Downloading https://www.cs.toronto.edu/~kriz/cifar-10-python.tar.gz to ./data/cifar-10-python.tar.gz
Extracting ./data/cifar-10-python.tar.gz to ./data
Files already downloaded and verified
import matplotlib.pyplot as plt
import numpy as np

def imshow(img):
    img = img / 2 + 0.5
    npimg = img.numpy()
    plt.imshow(np.transpose(npimg, (1, 2, 0)))  # 原始数据是PILimage,BGR格式,plot只能显示RGB格式,必须要转置
    plt.show()

# 用迭代器来访问数据,一次访问的数据量是一个batch
dataiter = iter(trainloader)
images, labels = dataiter.next()

imshow(torchvision.utils.make_grid(images))  # make_grid用于给图像加上边框
print(' '.join('%5s' % classes[labels[j]] for j in range(4)))
horse   car   dog plane

步骤2 定义神经网络

前面的章节我们已经定义过神经网络了,直接将代码复用,修改为输入3通道即可。

import torch
import torch.nn as nn
import torch.nn.functional as F  # nn.functional提供了各种激励函数

class Net(nn.Module):
    
    def __init__(self):
        super(Net, self).__init__()
        # 这里将输入通道改为3
        self.conv1 = nn.Conv2d(3, 6, 5)
        self.conv2 = nn.Conv2d(6, 16, 5)
        
        self.fc1 = nn.Linear(16 * 5 * 5, 120)
        self.fc2 = nn.Linear(120, 84)
        self.fc3 = nn.Linear(84, 10)
    
    def forward(self, x):
        x = F.max_pool2d(F.relu(self.conv1(x)), (2, 2))
        x = F.max_pool2d(F.relu(self.conv2(x)), 2)
        
        x = x.view(-1, 16 * 5 * 5)
        x = F.relu(self.fc1(x))
        x = F.relu(self.fc2(x))
        x = self.fc3(x)
        
        return x

net = Net()

步骤3 误差计算和参数更新

import torch.optim as optim

criterion = nn.CrossEntropyLoss()
optimizer = optim.SGD(net.parameters(), lr=0.001, momentum=0.9)  # momentum表示动量, 一般设为0.9,带动量的梯度下降法收敛更快

步骤4 训练神经网络

for epoch in range(2):  # epoch表示在整个数据集上循环训练的次数
    
    running_loss = 0.0
    for i, data in enumerate(trainloader, 0):  #enumerate()将会给可迭代对象的元素标上序号,返回(序号, 元素)
        # 这里的data是以batch为单位的
        inputs, labels = data  # data的特征和标签分开
        
        # 清空梯度
        optimizer.zero_grad()
        
        # 处理输入、计算误差、更新权重
        outputs = net(inputs)
        loss = criterion(outputs, labels)
        loss.backward()
        optimizer.step()
        
        # 做一些统计
        running_loss += loss.item()  # loss是 1x1的Tenor,可以用item直接访问数据
        if i % 2000 == 1999:  # 每2000batch输出一次
            print('[%d, %5d] loss: %.3f' % (epoch + 1, i + 1, running_loss / 2000))
            running_loss = 0.0

print('Finished Training.')
[1,  2000] loss: 2.271
[1,  4000] loss: 1.946
[1,  6000] loss: 1.725
[1,  8000] loss: 1.598
[1, 10000] loss: 1.535
[1, 12000] loss: 1.477
[2,  2000] loss: 1.411
[2,  4000] loss: 1.389
[2,  6000] loss: 1.359
[2,  8000] loss: 1.340
[2, 10000] loss: 1.307
[2, 12000] loss: 1.280
Finished Training.

训练完成后,要记得保存训练好的模型:

PATH = './cifar_net.pth'
torch.save(net.state_dict(), PATH)

步骤5 测试神经网络

我们已经用数据集对神经网络训练了2遍,接下来要检验一下神经网络是否学到了东西。

检验的方法就是让神经网络再产生一些输出,并且和它们的标签做比对。

首先我们来看一组图片的标签:

dataiter = iter(testloader)
images, labels = dataiter.next()

imshow(torchvision.utils.make_grid(images))
print('GroundTruth: ', ' '.join('%5s' % classes[labels[j]] for j in range(4)))
GroundTruth:    cat  ship  ship plane

接下来我们导入保存好的模型,看看模型认为这些图片是什么。模型的输出是图片的“能量”,能量共有10个值,分别表示这场图片属于对应类别的可能性,能量越大,代表我们的分类器认为图片越属于一个类。

net = Net()
net.load_state_dict(torch.load(PATH))

outputs = net(images)
# torch.max不仅可以返回最大值,还可以返回最大值的索引(第二个返回值),我们不需要知道能量的具体值,只需要知道图片归属哪一类即可,最大能量对应的索引即是它被归为的类
_, predicted = torch.max(outputs, 1)  

print('Predicted: ', ' '.join('%5s' % classes[predicted[j]] for j in range(4)))
Predicted:    cat  ship  ship  ship

结果还算不错,接下来我们把网络应用到完整数据集上试一试:

correct = 0 
total = 0
with torch.no_grad():
    for data in testloader:
        images,labels = data
        outputs = net(images)
        _, predicted = torch.max(outputs.data, 1)
        total += labels.size(0)
        correct += (predicted == labels).sum().item()


print('Accuracy of the network on the 10000 test images: %d %%' % (100 * correct / total))     
Accuracy of the network on the 10000 test images: 55 %

再按类别做一次统计,看一看我们的网络的优势和短板是什么:

class_correct = list(0. for i in range(10))
class_total = list(0. for i in range(10))
with torch.no_grad():
    for data in testloader:
        images, labels = data
        outputs = net(images)
        _, predicted = torch.max(outputs, 1)
        c = (predicted == labels).squeeze()
        for i in range(4):
            label = labels[i]
            class_correct[label] += c[i].item()
            class_total[label] += 1


for i in range(10):
    print('Accuracy of %5s : %2d %%' % (
        classes[i], 100 * class_correct[i] / class_total[i]))
Accuracy of plane : 70 %
Accuracy of   car : 67 %
Accuracy of  bird : 34 %
Accuracy of   cat : 43 %
Accuracy of  deer : 52 %
Accuracy of   dog : 52 %
Accuracy of  frog : 67 %
Accuracy of horse : 58 %
Accuracy of  ship : 60 %
Accuracy of truck : 51 %

Training on GPU

在GPU上进行训练也非常简单,怎么把Tensor转到GPU,就怎么把网络转到GPU:

device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")

print(device)
cuda:0

接下来我们直接使用net.to(device)即可把网络迁移到GPU上,程序会自动识别所有的参数,将他们转化为CUDA Tensor。

需要注意的是,我们必须把输入的数据和标签也都迁移至GPU:

inputs, labels = data[0].to(device), data[1].to(device)

至此,Pytorch tutorial篇已经完结,官方原版第5篇教程Optional: Data Parallelism为可选部分,不再另行翻译。

最后编辑于
©著作权归作者所有,转载或内容合作请联系作者
  • 序言:七十年代末,一起剥皮案震惊了整个滨河市,随后出现的几起案子,更是在滨河造成了极大的恐慌,老刑警刘岩,带你破解...
    沈念sama阅读 199,636评论 5 468
  • 序言:滨河连续发生了三起死亡事件,死亡现场离奇诡异,居然都是意外死亡,警方通过查阅死者的电脑和手机,发现死者居然都...
    沈念sama阅读 83,890评论 2 376
  • 文/潘晓璐 我一进店门,熙熙楼的掌柜王于贵愁眉苦脸地迎上来,“玉大人,你说我怎么就摊上这事。” “怎么了?”我有些...
    开封第一讲书人阅读 146,680评论 0 330
  • 文/不坏的土叔 我叫张陵,是天一观的道长。 经常有香客问我,道长,这世上最难降的妖魔是什么? 我笑而不...
    开封第一讲书人阅读 53,766评论 1 271
  • 正文 为了忘掉前任,我火速办了婚礼,结果婚礼上,老公的妹妹穿的比我还像新娘。我一直安慰自己,他们只是感情好,可当我...
    茶点故事阅读 62,665评论 5 359
  • 文/花漫 我一把揭开白布。 她就那样静静地躺着,像睡着了一般。 火红的嫁衣衬着肌肤如雪。 梳的纹丝不乱的头发上,一...
    开封第一讲书人阅读 48,045评论 1 276
  • 那天,我揣着相机与录音,去河边找鬼。 笑死,一个胖子当着我的面吹牛,可吹牛的内容都是我干的。 我是一名探鬼主播,决...
    沈念sama阅读 37,515评论 3 390
  • 文/苍兰香墨 我猛地睁开眼,长吁一口气:“原来是场噩梦啊……” “哼!你这毒妇竟也来了?” 一声冷哼从身侧响起,我...
    开封第一讲书人阅读 36,182评论 0 254
  • 序言:老挝万荣一对情侣失踪,失踪者是张志新(化名)和其女友刘颖,没想到半个月后,有当地人在树林里发现了一具尸体,经...
    沈念sama阅读 40,334评论 1 294
  • 正文 独居荒郊野岭守林人离奇死亡,尸身上长有42处带血的脓包…… 初始之章·张勋 以下内容为张勋视角 年9月15日...
    茶点故事阅读 35,274评论 2 317
  • 正文 我和宋清朗相恋三年,在试婚纱的时候发现自己被绿了。 大学时的朋友给我发了我未婚夫和他白月光在一起吃饭的照片。...
    茶点故事阅读 37,319评论 1 329
  • 序言:一个原本活蹦乱跳的男人离奇死亡,死状恐怖,灵堂内的尸体忽然破棺而出,到底是诈尸还是另有隐情,我是刑警宁泽,带...
    沈念sama阅读 33,002评论 3 315
  • 正文 年R本政府宣布,位于F岛的核电站,受9级特大地震影响,放射性物质发生泄漏。R本人自食恶果不足惜,却给世界环境...
    茶点故事阅读 38,599评论 3 303
  • 文/蒙蒙 一、第九天 我趴在偏房一处隐蔽的房顶上张望。 院中可真热闹,春花似锦、人声如沸。这庄子的主人今日做“春日...
    开封第一讲书人阅读 29,675评论 0 19
  • 文/苍兰香墨 我抬头看了看天上的太阳。三九已至,却和暖如春,着一层夹袄步出监牢的瞬间,已是汗流浃背。 一阵脚步声响...
    开封第一讲书人阅读 30,917评论 1 255
  • 我被黑心中介骗来泰国打工, 没想到刚下飞机就差点儿被人妖公主榨干…… 1. 我叫王不留,地道东北人。 一个月前我还...
    沈念sama阅读 42,309评论 2 345
  • 正文 我出身青楼,却偏偏与公主长得像,于是被迫代替她去往敌国和亲。 传闻我的和亲对象是个残疾皇子,可洞房花烛夜当晚...
    茶点故事阅读 41,885评论 2 341

推荐阅读更多精彩内容