尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

深耕网站建设与建站编程的一线实战洞察。

10-CNN案例-图像分类

10-CNN案例-图像分类 需求说明实现图像分类。数据说明CIFAR-10数据集5万张训练图像、1万张测试图像、10个类别、每个类别有6k个图像图像大小32×32×3。下图列举了10个类每一类随机展示了10张图片建模import torch import torch.nn as nn from torchvision.datasets import CIFAR10 from torchvision.transforms import ToTensor # pip install torchvision -i https://mirrors.aliyun.com/pypi/simple/ import torch.optim as optim from torch.utils.data import DataLoader import time import matplotlib.pyplot as plt from torchsummary import summary # 每批次样本数 BATCH_SIZE 8加载数据# todo 1: 准备数据 def create_dataset(): train_dataset CIFAR10(root./data, trainTrue, transformToTensor(), downloadTrue) test_dataset CIFAR10(root./data, trainFalse, transformToTensor(), downloadTrue) return train_dataset, test_dataset搭建神经网络# todo 2: 搭建神经网络 输入形状: 32x32 第一个卷积层输入 3 个 Channel, 输出 6 个 Channel, Kernel Size 为: 3x3 第一个池化层输入 30x30, 输出 15x15, Kernel Size 为: 2x2, Stride 为: 2 第二个卷积层输入 6 个 Channel, 输出 16 个 Channel, Kernel Size 为 3x3 第二个池化层输入 13x13, 输出 6x6, Kernel Size 为: 2x2, Stride 为: 2 第一个全连接层输入 576 维, 输出 120 维 第二个全连接层输入 120 维, 输出 84 维 最后的输出层输入 84 维, 输出 10 维 class ImageModel(nn.Module): # 1. 初始化父类成员 def __init__(self): # 1.1 父类成员 super().__init__() # 1.2 搭建神经网络 # 卷积层和池化层 self.conv1 nn.Conv2d(3, 6, 3, 1, 0) self.pool1 nn.MaxPool2d(2, 2,0) self.conv2 nn.Conv2d(6, 16, 3, 1, 0) self.pool2 nn.MaxPool2d(2, 2,0) # 全连接层 self.linear1 nn.Linear(576, 120) self.linear2 nn.Linear(120, 84) # 输出层 self.output nn.Linear(84, 10) # 2. 前向传播 def forward(self, x): # 第1层卷积加权求和-激励层激活函数-池化层降维 x self.pool1(torch.relu(self.conv1(x))) # 第2层卷积加权求和-激励层激活函数-池化层降维 x self.pool2(torch.relu(self.conv2(x))) # 拉平数据全连接层只能处理二位数据所以要将数据进行拉平 x x.reshape(x.size(0),-1) # print(f全连接层输入形状{x.shape}) # 第3层全连接层加权求和激励层激活函数 x torch.relu(self.linear1(x)) # 第4层全连接层加权求和激励层激活函数 x torch.relu(self.linear2(x)) # 输出层全连接层加权求和多分类这里可以不用softmax因为后续用多分类交叉熵损失函数CrossEntropyLoss() x self.output(x) return x模型训练# todo 3: 模型训练 def train_model(train_dataset): # 1. 创建数据加载器 dataloader DataLoader(train_dataset, batch_sizeBATCH_SIZE, shuffleTrue) # 2. 创建模型对象 model ImageModel() # 3. 创建损失函数对象 criterion nn.CrossEntropyLoss() # 多分类交叉熵损失函数 softmax激活函数 损失计算 # 4. 创建优化器对象 optimizer optim.Adam(model.parameters(), lr1e-3) # 5. 循环遍历epoch开始每轮的训练 # 5.1 定义训练轮数 epochs 10 # 5.2 遍历完成每轮所有批次的训练 for epoch in range(epochs): # 5.2.1 记录总损失、总样本数据量、预测正确的样本数据量、训练开始时间 total_loss, total_samples, total_correct, start 0.0, 0, 0, time.time() # 5.2.2 遍历数据加载器获取每批次数据 for x,y in dataloader: # 5.2.2.1 切换模型为训练模式 model.train() # 5.2.2.2 模型预测 y_pred model(x) # 5.2.2.3 计算损失 loss criterion(y_pred, y) # 5.2.2.4 梯度清零 反向传播 优化器更新参数 optimizer.zero_grad() loss.backward() optimizer.step() # 5.2.2.5 统计预测正确的样本个数 # print(torch.argmax(y_pred, dim-1)) # -1这里是指最后一维这里是行 total_correct (torch.argmax(y_pred, dim-1) y).sum() # 5.2.2.6 统计当前批次的总损失 total_loss loss.item() * len(y) # 5.2.2.7 统计当前批次的总样本数 total_samples len(y) # break # 只训练1个批次提高训练效率用于测试实际训练不能这么写 # print(**50) # 5.2.3 打印每轮的训练结果走这里表示1轮训练已经完成 print(fepoch: {epoch1}, loss: {total_loss/total_samples:.5f}, accuracy: {total_correct/total_samples:.2f}, time: {time.time() - start:.2f}s) # break # 只训练1轮提高训练效率用于测试实际训练不能这么写 # 5.3 保存模型 torch.save(model.state_dict(), ./model/image_model.pth)模型预测# todo 4: 模型测试 def test_model(test_dataset): # 1. 创建测试机数据加载器 dataloader DataLoader(test_dataset, batch_sizeBATCH_SIZE, shuffleTrue) # 2. 创建模型对象 model ImageModel() # 3. 加载模型参数 model.load_state_dict(torch.load(./model/image_model.pth)) # 4. 统计预测正确的样本个数、总样本个数 total_correct, total_samples 0,0 # 5. 遍历数据加载器获取每批次的数据 for x,y in dataloader: # 5.1 切换模型模式 model.eval() # 5.2 模型预测 y_pred model(x) # 5.3 argmax()模拟softmax y_pred torch.argmax(y_pred, dim-1) # 5.4 统计正确样本数 total_correct (y_pred y).sum() # 5.5 统计总样本个数 total_samples len(y) # 6. print(facc: {total_correct/total_samples:.2f})优化思路增加卷积核输出通道数增加全连接层的参数量调整学习率调整优化方法修改激活函数……测试if __name__ __main__: train_dataset, test_dataset create_dataset() # print(f训练集{train_dataset.data.shape}) # print(f测试集{test_dataset.data.shape}) # print(f训练集类别{train_dataset.class_to_idx}) # # # 图像展示 # plt.figure(figsize(2, 2)) # plt.imshow(train_dataset.data[11]) # plt.title(train_dataset.targets[11]) # plt.show() # 2. 搭建神经网络 # model ImageModel() # 查看模型参数 # 参1模型参2输入数据形状CHW参3批次大小 # summary(model, (3, 32, 32), batch_sizeBATCH_SIZE) # 3. 模型训练 # train_model(train_dataset) # 4. 模型测试 test_model(test_dataset)
返回列表