PyTorch食物图像分类实战:从数据集制作到CNN模型训练全流程详解

前言

很多刚入门深度学习的同学都会困惑:自己本地的图片数据集,到底怎么一步步变成模型能训练的数据?CNN每一层的维度该怎么算?训练和测试模式到底有什么区别?这篇博客就以食物图像识别为例,从零带大家走一遍完整流程,代码拆成小段讲,把原理说透。


一、数据集准备:生成标签索引文件

图像分类任务的数据集,通常都是「按类别分文件夹」的目录结构:

复制代码
food_dataset/
├── train/
│   ├── 类别1/
│   │   ├── 001.jpg
│   │   └── 002.jpg
│   ├── 类别2/
│   └── ...
└── test/
    ├── 类别1/
    └── ...

PyTorch不能直接读取这种文件夹结构,我们需要先生成一个文本索引文件,每一行存「图片完整路径 + 空格 + 数字类别标签」,后续读取数据会非常方便。

1. 逐段代码讲解

首先导入操作系统模块,用来处理文件路径:

复制代码
import os

定义一个通用处理函数,两个参数分别是数据集根目录、要处理的子目录名(train或test):

复制代码
def train_test_file(root, dir):
    # 创建并打开要写入的txt文件,比如传入train就生成train.txt
    file_txt = open(dir + '.txt', 'w')
    # 拼接出当前要遍历的完整目录路径
    path = os.path.join(root, dir)

核心是os.walk,它会递归遍历目录,每一轮返回三个值:

  • roots:当前遍历到的文件夹的完整路径

  • directories:当前文件夹下所有子文件夹的名字列表

  • files:当前文件夹下所有文件的名字列表

    复制代码
      for roots, directories, files in os.walk(path):
          # 如果还有子文件夹,说明当前在类别目录的上一层
          # 先把所有类别文件夹名存下来,后续用下标当标签
          if len(directories) != 0:
              dirs = directories
          else:
              # 没有子文件夹了,说明已经进入某个类别文件夹内部
              # 拆分路径,取出当前类别文件夹的名字
              now_dir = roots.split('\\')
              # 遍历这个类别下的所有图片文件
              for file in files:
                  # 拼接出单张图片的完整路径
                  path_1 = os.path.join(roots, file)
                  print(path_1)
                  # 写入txt:图片路径 + 空格 + 类别索引 + 换行
                  # dirs.index(now_dir[-1]) 就是找到当前类别在类别列表里的序号,作为标签
                  file_txt.write(path_1 + ' ' + str(dirs.index(now_dir[-1])) + '\n')
      file_txt.close()

最后调用函数,分别生成训练集和测试集的索引文件:

复制代码
root = r'.\food_dataset'
train_dir = 'train'
test_dir = 'test'
train_test_file(root, train_dir)
train_test_file(root, test_dir)

运行完成后,当前目录下会出现train.txttest.txt,打开就能看到每一行都是「图片路径 标签」的格式。标签是从0开始的数字,和类别文件夹的顺序一一对应。


二、搞懂Dataset的核心:两个魔法方法

在PyTorch里自定义数据集,必须继承Dataset类,并且实现__len____getitem__两个方法。很多同学刚接触觉得很抽象,我们先用一个最简单的例子搞懂它们的作用。

1. 入门小例子

我们自己写一个类,实现这两个方法,看看效果:

复制代码
class USE_getitem():
    # 初始化方法:创建对象的时候自动执行,这里我们存一个字符串
    def __init__(self, text):
        self.text = text
    
    # 实现__getitem__:对象就可以用 [下标] 的方式取值
    def __getitem__(self, index):
        # 这里我们返回对应位置字符的大写形式
        result = self.text[index].upper()
        return result
    
    # 实现__len__:就可以用 len(对象) 获取长度
    def __len__(self):
        return len(self.text)

我们实例化这个类,测试一下效果:

复制代码
p = USE_getitem("pytorch")

# 两种写法效果完全一样,p[1]会自动调用__getitem__方法
print(p.__getitem__(1))  # 输出 Y
print(p[1])              # 输出 Y

# 两种获取长度的方式同理,len(p)会自动调用__len__方法
print(p.__len__())       # 输出 7
print(len(p))            # 输出 7

print(p[0], p[1])        # 输出 P Y

2. 原理总结

这两个都是Python的魔法方法:

  • 实现__getitem__,自定义对象就能像列表、字符串一样,用[索引]获取元素;
  • 实现__len__,就能用len()函数获取对象的长度。

PyTorch的Dataset就是基于这个机制设计的。只要我们的类继承了Dataset,并且正确实现这两个方法,PyTorch就能把它识别为合法的数据集,后续才能用DataLoader批量加载。


三、自定义食物数据集类

搞懂了基础原理,我们来写真正能用的数据集类,用来读取刚才生成的txt标签文件。

1. 导入依赖工具

复制代码
import torch
import torch.nn as nn
from torch.utils.data import Dataset, DataLoader
import numpy as np
from PIL import Image
from torchvision import transforms

2. 定义数据预处理

图片不能直接扔进模型训练,需要先做统一的预处理。transforms.Compose可以把多个变换操作按顺序组合起来。

复制代码
data_transforms = {
    # 训练集预处理
    'trainda':
        transforms.Compose([
            transforms.Resize([256, 256]),  # 把所有图片缩放到256×256的统一尺寸
            transforms.ToTensor(),          # 转成Tensor张量:通道顺序从HWC变CHW,像素值归一化到0~1
        ]),
    # 测试集预处理,缩放规则和训练集保持一致
    'valid':
        transforms.Compose([
            transforms.Resize([256, 256]),
            transforms.ToTensor(),
        ]),
}

补充说明:ToTensor()做了两件关键的事:

  1. 把PIL图像的「高×宽×通道」(HWC)格式,转换成卷积层要求的「通道×高×宽」(CHW)格式;
  2. 把0-255的整数像素值除以255,归一化到0-1的浮点数区间。

3. 编写数据集类

我们的类继承Dataset,一共实现三个方法:初始化、获取长度、根据索引取样本。

首先是初始化方法,负责读取txt文件,把所有图片路径和标签存到列表里:

复制代码
class food_dataset(Dataset):
    def __init__(self, file_path, transform=None):
        self.file_path = file_path
        self.imgs = []     # 存放所有图片的路径
        self.labels = []   # 存放每张图片对应的标签
        self.transform = transform
        
        # 打开txt文件,逐行读取
        with open(self.file_path) as f:
            # 每一行按空格拆分,得到图片路径和标签
            samples = [x.strip().split(' ') for x in f.readlines()]
            for img_path, label in samples:
                self.imgs.append(img_path)
                self.labels.append(label)

__len__方法很简单,返回数据集总共有多少张图片:

复制代码
    def __len__(self):
        return len(self.imgs)

最核心的__getitem__方法,根据索引返回一张处理好的图片和对应的标签张量:

复制代码
    def __getitem__(self, idx):
        # 1. 根据路径读取图片,此时是PIL图像格式
        image = Image.open(self.imgs[idx])
        
        # 2. 如果传入了预处理,就对图片执行变换
        if self.transform:
            image = self.transform(image)
        
        # 3. 把标签转成int64类型的张量,和模型输出格式匹配
        label = self.labels[idx]
        label = torch.from_numpy(np.array(label, dtype=np.int64))
        
        # 返回 图片张量 + 标签张量
        return image, label

4. 实例化数据集对象

复制代码
# 训练集:传入train.txt路径,使用训练集预处理
training_data = food_dataset(file_path=r'./train.txt', transform=data_transforms['trainda'])
# 测试集:传入test.txt路径,使用测试集预处理
test_data = food_dataset(file_path=r'./test.txt', transform=data_transforms['valid'])

四、DataLoader:批量加载数据

Dataset只能一张一张地取样本,实际训练时我们都是一批一批地喂给模型,这就需要DataLoader来做打包、打乱、并行加载等工作。

复制代码
# 训练集加载器:每批64张图,打乱顺序
train_dataloader = DataLoader(training_data, batch_size=64, shuffle=True)
# 测试集加载器
test_dataloader = DataLoader(test_data, batch_size=64, shuffle=True)

几个关键参数说明:

  • batch_size:每个批次包含多少张图片,根据显存大小调整,越大训练越快但越占显存;
  • shuffle:是否打乱数据顺序,训练集建议开启,避免模型记住数据顺序;
  • num_workers:用多少个进程加载数据,Windows环境下如果报错可以设为0。

五、搭建CNN卷积神经网络

接下来我们搭建一个基础的卷积神经网络,完成20类食物图像分类。这里我会带着大家一步步算每一层的输出维度,彻底搞懂维度变化。

1. 先记卷积尺寸公式

卷积层输出的尺寸计算公式:

复制代码
输出尺寸 = (输入尺寸 + 2×padding - kernel_size) / stride + 1

如果是kernel_size=2的最大池化且步长为2,输出尺寸直接是输入的一半。

2. 网络结构逐段拆解

我们的网络包含3个卷积模块,最后接全连接层输出分类结果。

复制代码
class CNN(nn.Module):
    def __init__(self):
        super(CNN, self).__init__()
        
        # 第一个卷积模块:卷积 + ReLU激活 + 最大池化
        self.conv1 = nn.Sequential(
            nn.Conv2d(
                in_channels=3,      # 输入通道数:RGB彩色图是3通道
                out_channels=16,    # 输出通道数:也就是卷积核的个数
                kernel_size=5,      # 卷积核大小 5×5
                stride=1,           # 卷积核移动步长
                padding=2,          # 边缘填充像素数
            ),
            nn.ReLU(),              # ReLU激活函数,引入非线性
            nn.MaxPool2d(kernel_size=2),  # 2×2最大池化,压缩尺寸
        )

维度计算(conv1) : 输入形状:3 × 256 × 256 卷积后尺寸:(256 + 2×2 - 5) / 1 + 1 = 256 → 形状 16 × 256 × 256 池化后尺寸:256 / 2 = 128 → 输出形状 16 × 128 × 128

继续第二个卷积模块:

复制代码
        # 第二个卷积模块:两个卷积 + ReLU + 最大池化
        self.conv2 = nn.Sequential(
            nn.Conv2d(16, 32, 5, 1, 2),  # 输入16通道,输出32通道
            nn.ReLU(),
            nn.Conv2d(32, 32, 5, 1, 2),  # 输入32通道,输出32通道
            nn.ReLU(),
            nn.MaxPool2d(2),
        )

维度计算(conv2) : 输入形状:16 × 128 × 128 两次卷积后尺寸保持128不变 → 形状 32 × 128 × 128 池化后尺寸:128 / 2 = 64 → 输出形状 32 × 64 × 64

第三个卷积模块:

复制代码
        # 第三个卷积模块:卷积 + ReLU
        self.conv3 = nn.Sequential(
            nn.Conv2d(32, 128, 5, 1, 2),  # 输入32通道,输出128通道
            nn.ReLU(),
        )

维度计算(conv3) : 输入形状:32 × 64 × 64 卷积后尺寸保持64不变 → 输出形状 128 × 64 × 64

最后是全连接层,需要先把二维特征图展平成一维向量,再输出20个类别的预测结果:

复制代码
        # 全连接层:输入维度 = 通道数 × 高 × 宽 = 128×64×64,输出20个类别
        self.out = nn.Linear(128 * 64 * 64, 20)

然后是前向传播函数forward,定义数据流过网络的顺序:

复制代码
    def forward(self, x):
        x = self.conv1(x)
        x = self.conv2(x)
        x = self.conv3(x)
        # 展平操作:把(batch_size, 通道, 高, 宽)变成(batch_size, 通道*高*宽)
        x = x.view(x.size(0), -1)
        output = self.out(x)
        return output

六、设备选择:用GPU加速训练

有显卡的话一定要用GPU训练,速度会比CPU快很多。我们写一段自动判断设备的代码,兼容N卡、苹果M系列芯片和CPU。

复制代码
# 优先级:cuda(N卡) > mps(苹果M系列) > CPU
device = "cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu"

# 把模型移动到对应设备上
model = CNN().to(device)
print(model)

运行后会打印出完整的网络结构,大家可以核对每一层的输入输出通道数。


七、编写训练函数

训练的核心流程就是:喂入一批数据 → 模型前向预测 → 计算损失 → 反向传播更新参数 → 循环。

复制代码
def train(dataloader, model, loss_fn, optimizer):
    # 切换到训练模式:启用Dropout、BatchNorm等训练专属层
    model.train()
    
    batch_size_num = 1  # 记录当前是第几个批次
    # 遍历每一批数据
    for X, y in dataloader:
        # 把数据和标签也移到设备上,必须和模型在同一个设备
        X, y = X.to(device), y.to(device)
        
        # 1. 前向传播,得到模型预测结果
        pred = model.forward(X)
        # 2. 计算预测值和真实标签的损失
        loss = loss_fn(pred, y)
        
        # 3. 梯度清零(非常重要!PyTorch默认梯度会累加)
        optimizer.zero_grad()
        # 4. 反向传播,计算每个参数的梯度
        loss.backward()
        # 5. 根据梯度更新网络参数
        optimizer.step()
        
        # 打印当前损失
        loss_value = loss.item()
        if batch_size_num % 1 == 0:
            print(f"loss: {loss_value:>7f}  [number:{batch_size_num}]")
        batch_size_num += 1

重点提醒:optimizer.zero_grad()这一步绝对不能忘!如果不清零,梯度会和上一批次累加,参数更新就完全错了。


八、编写测试函数

测试的时候模型参数是固定的,不需要计算梯度,这样既能省显存,速度也更快。

复制代码
def test(dataloader, model, loss_fn):
    size = len(dataloader.dataset)  # 测试集总样本数
    num_batches = len(dataloader)   # 总批次数量
    # 切换到评估模式:固定参数,关闭Dropout等
    model.eval()
    test_loss, correct = 0, 0
    
    # 关闭梯度计算上下文
    with torch.no_grad():
        for X, y in dataloader:
            X, y = X.to(device), y.to(device)
            pred = model.forward(X)
            
            # 累加每一批的损失
            test_loss += loss_fn(pred, y).item()
            # 统计预测正确的数量:取概率最大的类别和真实标签比较
            correct += (pred.argmax(1) == y).type(torch.float).sum().item()
    
    # 计算平均损失和整体准确率
    test_loss /= num_batches
    correct /= size
    print(f"Test result: \n Accuracy: {(100*correct)}%, Avg loss: {test_loss}")

补充:pred.argmax(1)表示在维度1(类别维度)上找最大值的索引,也就是模型认为最可能的类别。


九、开始训练

准备工作都做完了,现在设置损失函数、优化器和训练轮数,正式开始训练。

复制代码
# 交叉熵损失函数,多分类任务的标准选择
loss_fn = nn.CrossEntropyLoss()
# Adam优化器,传入模型参数和学习率
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)

# 总共训练15轮
epochs = 15
for t in range(epochs):
    print(f"Epoch {t+1}\n-------------------------------")
    train(train_dataloader, model, loss_fn, optimizer)

print("Done!")
# 训练全部结束后,用测试集评估最终效果
test(test_dataloader, model, loss_fn)

运行之后就能看到每一轮的loss逐步下降,训练完成后会输出测试集的准确率。


十、最后总结

到这里,一个完整的食物图像分类项目就从零跑通了。整个流程可以概括为:制作数据集索引文件 → 自定义Dataset类 → DataLoader批量加载 → 搭建CNN网络 → 训练与评估。

这只是一个基础入门版本,想要提升效果还有很多优化方向:

  • 加入随机翻转、旋转、裁剪等数据增强,缓解过拟合
  • 调整学习率,加入学习率衰减策略
  • 更换更深的网络结构,比如ResNet、MobileNet等
  • 加入验证集,使用早停策略防止过拟合

大家可以自己动手修改尝试,深度学习多跑多调参,慢慢就有手感了。

相关推荐
明志数科1 小时前
从300克Ego头环看第一人称数据采集趋势:设备轻量化之后,场景端壁垒在哪
数码相机·学习
笨鸟先飞的橘猫2 小时前
系统设计第十七天决策卡
学习·游戏
早睡早起身体好1233 小时前
用 XGrammar 约束大模型工具调用:解决参数为空的问题
android·人工智能·神经网络·机器学习·自然语言处理·vllm
Sunshing153 小时前
Lyapunov方程系统本身稳定性判定与镇定性判定
笔记·学习
王红臣同学3 小时前
Microduck 强化学习源码拆解:一只 800g 的机器鸭怎么学会走路
人工智能·机器学习·ai
nagualky1233 小时前
AI Agent 别只返回“已完成”:用五字段验收契约定义任务终点
人工智能·机器学习·软件工程
2601_962301013 小时前
深入探索 TensorFlow 2.0:Python 中的强大深度学习框架
深度学习·分布式训练·keras·tensorflow2.0·eagerexecution
方方洛4 小时前
vllm教程-19-模型架构基础
人工智能·深度学习·llm
sunoo-2295 小时前
51单片机学习Day2:数码管动态扫描、定时器与中断系统深度总
单片机·嵌入式硬件·学习·51单片机