学习目的:学会构建CNN网络

一、 前期准备

关于环境

  • 语言环境:Python3.13
  • 编译器:vsCode
  • 深度学习环境:torch==2.11.0+cu130;torchvision==0.26.0+cu130

关于CIFAR10数据集:

        CIFAR10数据集是一批包含 10 个类别(飞机、汽车、鸟、猫、鹿、狗、青蛙、马、船、卡车)

1.设置GPU

设置分析环境:

import torch
import torch.nn as nn
import matplotlib.pyplot as plt
import torchvision

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

device

2. 导入数据

使用dataset下载CIFAR10数据集,并划分好训练集与测试集

使用dataloader加载数据,并设置好基本的batch_size

train_ds = torchvision.datasets.CIFAR10('data', 
                                      train=True, 
                                      transform=torchvision.transforms.ToTensor(), # 将数据类型转化为Tensor
                                      download=True)

test_ds  = torchvision.datasets.CIFAR10('data', 
                                      train=False, 
                                      transform=torchvision.transforms.ToTensor(), # 将数据类型转化为Tensor
                                      download=True)
batch_size = 32

train_dl = torch.utils.data.DataLoader(train_ds, 
                                       batch_size=batch_size, 
                                       shuffle=True)

test_dl  = torch.utils.data.DataLoader(test_ds, 
                                       batch_size=batch_size)

imgs, labels = next(iter(train_dl))
imgs.shape

代码运行结果:

运行结果解读:

 torch.Size([32, 3, 32, 32]),表示 32 张图片,每张是 3 通道(RGB),高 32 像素,宽 32 像素。显然这里输出的通道数和上周不同,上周由于识别的是黑白图片经转换后只有灰度值因此通道数为1.

3. 数据可视化

import numpy as np

# 指定图片大小,图像大小为20宽、5高的绘图(单位为英寸inch)
plt.figure(figsize=(20, 5)) 
for i, imgs in enumerate(imgs[:20]):
    # 进行轴变换
    npimg = imgs.numpy().transpose((1, 2, 0))
    # 将整个figure分成2行10列,绘制第i+1个子图。
    plt.subplot(2, 10, i+1)
    plt.imshow(npimg, cmap=plt.cm.binary)
    plt.axis('off')

代码理解:

transpose((1, 2, 0))详解:

  • 作用是对NumPy数组进行轴变换,transpose函数的参数是一个元组,定义了新轴的顺序。原始PyTorch张量通常是以(C, H, W)的格式存储的,其中:
    • C是通道数(例如,RGB图像有3个通道)。
    • H是图像的高度。
    • W是图像的宽度。
  • transpose((1, 2, 0))将轴的顺序从(C, H, W)转换为(H, W, C),这使得数据格式更适合可视化和处理。

例子

假设有一个三维数组 arr,形状为 (3, 32, 32),含义是:

  • 维度 0:颜色通道(C=3 表示 RGB)

  • 维度 1:图像高度(H=32)

  • 维度 2:图像宽度(W=32)

执行 arr.transpose((1, 2, 0)) 后,新数组的形状变为 (32, 32, 3),维度含义变为:

  • 新维度 0 = 原维度 1(高度)

  • 新维度 1 = 原维度 2(宽度)

  • 新维度 2 = 原维度 0(通道)

简单记忆:原来顺序是 (C, H, W),想要变成 (H, W, C),所以把索引 1(H)放到第 0 位,索引 2(W)放到第 1 位,索引 0(C)放到第 2 位,即 (1, 2, 0)

根据上一步代码运行结果torch.Size([32, 3, 32, 32]),可知如果我们取出其中一张图,默认储存格式是 (3, 32, 32),经过transpose((1, 2, 0)) 变成 (32, 32, 3)之后plt.imshow(npimg) 就能正确显示彩色图片了。

为什么第一周黑白图像没有用到这个代码呢?

因为黑白(灰度)图片通常只有一个通道,形状为 (1, H, W) 或直接 (H, W)Matplotlib 的 imshow 可以直接接受二维数组 (H, W) 作为灰度图,不需要转换成三维格式。

所以,黑白图片中没有用到是因为灰度图本来就是二维的,不需要通道置换;而彩色图必须是 (H, W, C) 才能正确显示 RGB 颜色。

代码运行结果:

二、构建简单的CNN网络

对于一般的CNN网络来说,都是由特征提取网络和分类网络构成,其中特征提取网络用于提取图片的特征,分类网络用于将图片进行分类。

import torch.nn.functional as F

num_classes = 10

class Model(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(3, 64, kernel_size=3)
        self.pool1 = nn.MaxPool2d(kernel_size=2)
        self.conv2 = nn.Conv2d(64, 64, kernel_size=3)
        self.pool2 = nn.MaxPool2d(kernel_size=2)
        self.conv3 = nn.Conv2d(64, 128, kernel_size=3)
        self.pool3 = nn.MaxPool2d(kernel_size=2)

        self.fc1 = nn.Linear(512, 256)
        self.fc2 = nn.Linear(256, num_classes)

    def forward(self, x):
        x = self.pool1(F.relu(self.conv1(x)))
        x = self.pool2(F.relu(self.conv2(x)))
        x = self.pool3(F.relu(self.conv3(x)))

        x = torch.flatten(x, start_dim=1)

        x = F.relu(self.fc1(x))
        x = self.fc2(x)

        return x

from torchinfo import summary

model = Model().to(device)

summary(model)

1. torch.nn.Conv2d()详解

函数原型:torch.nn.Conv2d(in_channels, out_channels, kernel_size, stride=1, padding=0, dilation=1, groups=1, bias=True, padding_mode='zeros', device=None, dtype=None)

关键参数说明

  • in_channels ( int ) – 输入图像中的通道数
  • out_channels ( int ) – 卷积产生的通道数
  • kernel_size ( int or tuple ) – 卷积核的大小
  • stride ( int or tuple , optional ) -- 卷积的步幅。默认值:1
  • padding ( int , tuple或str , optional ) – 添加到输入的所有四个边的填充。默认值:0
  • dilation (int or tuple, optional) - 扩张操作:控制kernel点(卷积核点)的间距,默认值:1。
  • groups(int,可选):将输入通道分组成多个子组,每个子组使用一组卷积核来处理。默认值为 1,表示不进行分组卷积。
  • padding_mode (字符串,可选) – 'zeros', 'reflect', 'replicate'或'circular'. 默认:'zeros'

2. torch.nn.Linear()详解

函数原型:torch.nn.Linear(in_features, out_features, bias=True, device=None, dtype=None)

关键参数说明

  • in_features:每个输入样本的大小
  • out_features:每个输出样本的大小

3. torch.nn.MaxPool2d()详解

函数原型:torch.nn.MaxPool2d(kernel_size, stride=None, padding=0, dilation=1, return_indices=False, ceil_mode=False)

关键参数说明

  • kernel_size:最大的窗口大小
  • stride:窗口的步幅,默认值为kernel_size
  • padding:填充值,默认为0
  • dilation:控制窗口中元素步幅的参数

手动推导过程:

参数量计算

卷积层参数量计算公式=k2×cin​×cout​+cout​

  • conv1(3×3×3)×64 + 64 = (27×64) + 64 = 1728 + 64 = 1792

  • conv2(3×3×64)×64 + 64 = (576×64) + 64 = 36864 + 64 = 36928 

  • conv3(3×3×64)×128 + 128 = (576×128) + 128 = 73728 + 128 = 73856

全连接层参数量计算

  • fc1(512×256) + 256 = 131072 + 256 = 131328 

  • fc2(256×10) + 10 = 2560 + 10 = 2570 

总参数:1792+36928+73856+131328+2570 = 246,474 

代码运行结果:

三、训练模型

1.超参数设置

loss_fn    = nn.CrossEntropyLoss()
learn_rate = 1e-2
opt        = torch.optim.SGD(model.parameters(),lr=learn_rate)

2. 编写训练函数

def train(dataloader, model, loss_fn, optimizer):
    size = len(dataloader.dataset)
    num_batches = len(dataloader)

    train_loss, train_acc = 0, 0

    for X, y in dataloader:
        X, y = X.to(device),y.to(device)

        pred = model(X)
        loss = loss_fn(pred, y)

        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

        train_acc  += (pred.argmax(1) == y).type(torch.float).sum().item()
        train_loss += loss.item()

    train_acc  /= size
    train_loss /= num_batches

    return train_acc, train_loss

3. 编写测试函数

def test (dataloader, model, loss_fn):
    size        = len(dataloader.dataset)  # 测试集的大小,一共10000张图片
    num_batches = len(dataloader)          # 批次数目,313(10000/32=312.5,向上取整)
    test_loss, test_acc = 0, 0
    
    # 当不进行训练时,停止梯度更新,节省计算内存消耗
    with torch.no_grad():
        for imgs, target in dataloader:
            imgs, target = imgs.to(device), target.to(device)
            
            # 计算loss
            target_pred = model(imgs)
            loss        = loss_fn(target_pred, target)
            
            test_loss += loss.item()
            test_acc  += (target_pred.argmax(1) == target).type(torch.float).sum().item()

    test_acc  /= size
    test_loss /= num_batches

    return test_acc, test_loss

4. 正式训练

epochs     = 10
train_loss = []
train_acc  = []
test_loss  = []
test_acc   = []

for epoch in range(epochs):
    model.train()
    epoch_train_acc, epoch_train_loss = train(train_dl, model, loss_fn, opt)
    
    model.eval()
    epoch_test_acc, epoch_test_loss = test(test_dl, model, loss_fn)
    
    train_acc.append(epoch_train_acc)
    train_loss.append(epoch_train_loss)
    test_acc.append(epoch_test_acc)
    test_loss.append(epoch_test_loss)
    
    template = ('Epoch:{:2d}, Train_acc:{:.1f}%, Train_loss:{:.3f}, Test_acc:{:.1f}%,Test_loss:{:.3f}')
    print(template.format(epoch+1, epoch_train_acc*100, epoch_train_loss, epoch_test_acc*100, epoch_test_loss))
print('Done')

代码运行结果:

四、 结果可视化

import matplotlib.pyplot as plt
#隐藏警告
import warnings
warnings.filterwarnings("ignore")               #忽略警告信息
plt.rcParams['font.sans-serif']    = ['SimHei'] # 用来正常显示中文标签
plt.rcParams['axes.unicode_minus'] = False      # 用来正常显示负号
plt.rcParams['figure.dpi']         = 100        #分辨率

from datetime import datetime
current_time = datetime.now() # 获取当前时间

epochs_range = range(epochs)

plt.figure(figsize=(12, 3))
plt.subplot(1, 2, 1)

plt.plot(epochs_range, train_acc, label='Training Accuracy')
plt.plot(epochs_range, test_acc, label='Test Accuracy')
plt.legend(loc='lower right')
plt.title('Training and Validation Accuracy')
plt.xlabel(current_time) 

plt.subplot(1, 2, 2)
plt.plot(epochs_range, train_loss, label='Training Loss')
plt.plot(epochs_range, test_loss, label='Test Loss')
plt.legend(loc='upper right')
plt.title('Training and Validation Loss')
plt.show()

代码运行结果:

五、感悟

这周进一步学习CNN网络,重点学习卷积层和池化层的整体的计算与推导。相对于黑白图像,彩色图像识别更难更复杂,且在前期数据准备的时候要注意转换。

在手敲代码过程中发现会有拼写错误,运行后报错,经常要逐行代码查错,需要进一步熟练基本代码,此外查报错技巧需要进一步掌握。

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐