【深度学习入门】用 PyTorch 从零搭建 CNN,完成 20 类食物图片识别


一、项目简介

我们的任务是:输入一张食物图片,输出它的类别。项目自己搭建了一个小型 CNN(卷积神经网络),完成 20 类食物 的分类:

八宝粥、哈密瓜、圣女果、巴旦木、板栗、汉堡、火龙果、炸鸡、瓜子、生肉、白萝卜、胡萝卜、草莓、菠萝、薯条、蛋、蛋挞、菠菜、骨肉相连、鸡翅。

技术栈与运行环境:

组件说明
PyTorch深度学习框架(含 torch.nntorch.utils.datatorch.optim
torchvision图像预处理工具 transforms
PIL (Pillow)读取图片
NumPy标签数组转换
设备自动选择 GPU (CUDA) / Apple 芯片 (MPS) / CPU

二、整体流程

先看一张流程图,理解全项目的骨架:

数据集文件夹
按类别分目录

生成 txt 标签文件
每行: 图片路径 标签

自定义 Dataset
实现 __len__ 与 __getitem__

DataLoader
分批 + 打乱加载

搭建 CNN 模型
卷积+池化+全连接

训练循环
前向 → 损失 → 反向传播 → 更新

测试集评估
计算准确率

保存最优模型
state_dict / TorchScript

单张图片预测
输出食物类别

各环节作用一览:

环节对应代码解决什么问题
生成 txt 标签os.walk 遍历目录把文件夹结构转成「路径 + 标签」文本
自定义 Datasetfood_dataset告诉框架如何按索引读一张图和它的标签
DataLoaderDataLoader(...)自动分 batch、打乱、并行加载
图像预处理transforms统一图片尺寸、转张量、归一化、数据增强
CNN 模型CNN(nn.Module)自动提取图片特征并完成分类
训练与评估train() / test()学习参数、验证效果
保存模型torch.save / torch.jit把训练成果持久化,便于部署
单图预测predict_image()用训练好的模型对新图片分类

三、数据准备:从文件夹生成 txt 标签文件

深度学习训练前,第一步是把数据整理成模型能用的形式。本项目没有直接使用现成的 torchvision.datasets.ImageFolder,而是自己写脚本生成 txt 标签文件,格式非常简单:

图片路径1 标签编号1
图片路径2 标签编号2
...

例如(这里用相对路径示意):

train/汉堡/001.jpg 5
train/草莓/012.jpg 12

其中标签编号就是类别文件夹在列表中的下标(从 0 开始)。

3.1 用 os.walk 遍历目录生成标签文件

使用 os.walk 遍历数据集目录生成 txt:

import os

def train_test_file(root, dir):
    file_txt = open(dir + '.txt', 'w')           # 创建 txt 文件
    path = os.path.join(root, dir)               # 拼接数据集路径
    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)
                file_txt.write(path_1 + ' ' + str(dirs.index(now_dir[-1])) + '\n')
    file_txt.close()

代码逻辑拆解:

  • os.walk(path) 会递归遍历目录,每次返回 (当前目录, 子目录列表, 文件列表)
  • 当遍历到根目录(存在子文件夹)时,记录下类别文件夹名列表 dirs
  • 当遍历到最底层的图片文件夹时,取文件夹名 now_dir[-1](即类别名),用 dirs.index(...) 得到它的编号,写入 txt。

四、核心概念:__getitem____len__

在写自定义 Dataset 之前,先理解两个 Python 魔术方法。用一个小例子做了演示:

class USE_getitem():
    def __init__(self, text):
        self.text = text

    def __getitem__(self, index):        # 支持下标取值:obj[index]
        result = self.text[index].upper()
        return result

    def __len__(self):                   # 支持 len(obj)
        return len(self.text)

p = USE_getitem('pytorch')
print(p[1])    # Y
print(len(p))  # 7

要点:

  • 实现了 __getitem__ 的对象,就可以像列表一样用 对象[下标] 取元素;
  • 实现了 __len__ 的对象,就可以用内置函数 len() 获取长度;
  • PyTorch 的 Dataset 正是依赖这两个方法:__len__ 告诉 DataLoader 一共有多少样本,__getitem__ 告诉 DataLoader 怎么按索引取第 idx 个样本。

五、自定义 Dataset 与 DataLoader

5.1 自定义 Dataset

torch.utils.data.Dataset 是所有数据集的基类,自定义数据集只要继承它并实现 __init____len____getitem__ 三个方法即可:

import torch
from torch.utils.data import Dataset, DataLoader
import numpy as np
from PIL import Image
from torchvision import transforms

class food_dataset(Dataset):
    def __init__(self, file_path, transform=None):
        self.file_path = file_path     # txt 标签文件路径
        self.imgs = []                 # 保存所有图片路径
        self.labels = []               # 保存所有图片标签
        self.transform = transform     # 图像预处理方法
        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)

    def __len__(self):
        return len(self.imgs)

    def __getitem__(self, idx):
        image = Image.open(self.imgs[idx])   # 按路径打开图片
        if self.transform:
            image = self.transform(image)    # 执行预处理 / 数据增强
        label = self.labels[idx]
        label = torch.from_numpy(np.array(label, dtype=np.int64))  # 转 int64 张量
        return image, label

逐段说明:

  • __init__:读取 txt,按行拆成「图片路径 + 标签」,分别存到两个列表;同时保存预处理方法;
  • __getitem__:用 PIL.Image.open 打开图片 → 做预处理 → 把字符串标签转成 int64 类型的 Tensor;
  • 标签为什么要转 int64?因为 CrossEntropyLoss 要求标签是 LongTensor(int64) 类型,否则会报类型错误。

5.2 DataLoader 批量加载

training_data = food_dataset(file_path='./train.txt', transform=data_transforms['train'])
test_data = food_dataset(file_path='./test.txt', transform=data_transforms['valid'])

train_dataloader = DataLoader(training_data, batch_size=64, shuffle=True)
test_dataloader = DataLoader(test_data, batch_size=64, shuffle=True)

DataLoader 两个常用参数:

参数作用
batch_size每个批次取多少张图(这里是 64 张)
shuffle每轮(epoch)是否打乱样本顺序,训练集通常 True,防止模型学到固定顺序

DataLoader 会自动把 __getitem__ 返回的单个样本堆叠成 batch:单张图片是 [C, H, W],一个 batch 就是 [batch_size, C, H, W] 的 Tensor。


六、图像预处理与数据增强(transforms)

6.1 基础预处理

基础的预处理只有两步:

data_transforms = {
    'train': transforms.Compose([
        transforms.Resize([256, 256]),  # 图片统一缩放到 256×256
        transforms.ToTensor(),          # PIL → Tensor,像素值归一化到 [0,1]
    ]),
    'valid': transforms.Compose([
        transforms.Resize([256, 256]),
        transforms.ToTensor(),
    ]),
}

两个操作的原理:

操作作用细节
Resize([256, 256])统一图片尺寸模型要求输入尺寸固定,所有图片必须缩放成一样大
ToTensor()转张量 + 归一化把 PIL 图片从 (H, W, C) 的 0~255 整数,变成 (C, H, W) 的 0~1 浮点数(通道顺序从 HWC 变为 CHW,这是 PyTorch 规定的格式)

6.2 进阶:Normalize 标准化

数据增强版本中还加入了归一化:

transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
  • 公式:(像素值 - mean) / std,对 R、G、B 三个通道分别做;
  • 这组数值是 ImageNet 数据集的均值与标准差,是图像任务最常用的归一化参数;
  • 归一化后数据分布更接近「均值 0、方差 1」,能帮助模型更快、更稳定地收敛

注意:Normalize 必须放在 ToTensor() 之后,因为它作用的对象是已经归一化到 [0,1] 的张量;同时训练集和验证集必须使用相同的归一化参数,否则数据分布不一致会导致评估失真。

6.3 数据增强(Data Augmentation)

数据增强 = 在训练时对图片做随机变换,生成"新的"样本。它不增加存储成本,却能成倍增加训练数据的多样性,是缓解过拟合最有效的手段之一。

数据增强版本中的完整训练集流水线:

data_transforms = {
    'train': transforms.Compose([
        transforms.Resize([256, 256]),              # 缩放
        transforms.RandomRotation(45),              # 随机旋转 ±45°
        transforms.CenterCrop(256),                 # 中心裁剪 256×256
        transforms.RandomHorizontalFlip(p=0.5),     # 50% 概率水平翻转
        transforms.RandomVerticalFlip(p=0.5),       # 50% 概率垂直翻转
        transforms.ColorJitter(brightness=0.2, contrast=0.1,
                               saturation=0.1, hue=0.1),  # 亮度/对比度/饱和度/色调扰动
        transforms.RandomGrayscale(p=0.1),          # 10% 概率转灰度(仍为 3 通道)
        transforms.ToTensor(),                      # 转张量
        transforms.Normalize([0.485, 0.456, 0.406],
                             [0.229, 0.224, 0.225]),  # 归一化
    ]),
    'valid': transforms.Compose([
        transforms.Resize([256, 256]),
        transforms.ToTensor(),
        transforms.Normalize([0.485, 0.456, 0.406],
                             [0.229, 0.224, 0.225]),
    ]),
}

各增强操作的含义:

操作参数效果
RandomRotation(45)45每张图随机旋转,角度在 ±45° 内
CenterCrop(256)256从中心裁剪出 256×256 区域
RandomHorizontalFlip(p=0.5)翻转概率 0.550% 概率左右镜像
RandomVerticalFlip(p=0.5)翻转概率 0.550% 概率上下镜像
ColorJitter(...)各分量扰动幅度随机调整亮度/对比度/饱和度/色调,模拟不同拍摄环境
RandomGrayscale(p=0.1)概率 0.110% 概率转为灰度图(输出仍是 3 通道)

关键原则:

  1. 增强只用于训练集。验证/测试集只用确定性预处理(Resize + ToTensor + Normalize),保证每次评估结果可复现;
  2. 增强操作都是随机的:同一张图每一轮训练看到的样子可能都不同,模型因此更鲁棒;
  3. 小细节:Resize(256) 后再 CenterCrop(256),裁剪结果与原图相同,这里起到的作用是统一尺寸;若想要更强增强效果,可以改用 RandomResizedCrop(随机位置 + 随机比例裁剪)。

七、CNN 模型搭建

7.1 完整模型代码

from torch import nn

class CNN(nn.Module):
    def __init__(self):
        super(CNN, self).__init__()
        # 卷积块 1:卷积 + 激活 + 最大池化
        self.conv1 = nn.Sequential(
            nn.Conv2d(in_channels=3, out_channels=16, kernel_size=5, stride=1, padding=2),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=2)   # 宽高减半
        )
        # 卷积块 2:两层卷积 + 激活 + 池化
        self.conv2 = nn.Sequential(
            nn.Conv2d(16, 32, 5, 1, 2),
            nn.ReLU(),
            nn.Conv2d(32, 32, 5, 1, 2),
            nn.ReLU(),
            nn.MaxPool2d(2)
        )
        # 卷积块 3:卷积 + 激活,不做池化
        self.conv3 = nn.Sequential(
            nn.Conv2d(32, 128, 5, 1, 2),
            nn.ReLU()
        )
        # 全连接层:展平后的特征映射 → 20 个类别
        self.out = nn.Linear(128 * 64 * 64, 20)

    def forward(self, x):
        x = self.conv1(x)
        x = self.conv2(x)
        x = self.conv3(x)
        x = x.view(x.size(0), -1)   # 展平,保留 batch 维度
        output = self.out(x)
        return output

7.2 Conv2d 参数详解

nn.Conv2d(in_channels, out_channels, kernel_size, stride, padding)

参数本项目取值含义
in_channels3(第一层)输入通道数,RGB 图为 3
out_channels16 / 32 / 128输出通道数,即用多少个卷积核,也决定特征图个数
kernel_size5卷积核大小 5×5
stride1卷积核每次滑动的步长
padding2边缘填充 2 圈 0,用于控制输出尺寸

特征图尺寸计算公式:

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

验证第一层:输入 256×256,kernel=5, padding=2, stride=1

(256 + 2×2 - 5) / 1 + 1 = 256

所以 padding=2 恰好让卷积不改变图片尺寸,只改变通道数。

7.3 网络结构尺寸变化

以输入 256×256×3 为例,数据在网络中的尺寸变化:

模块操作输出尺寸(H×W×C)
输入原始图片256×256×3
conv1Conv(3→16,5,1,2) + ReLU256×256×16
conv1MaxPool2d(2)128×128×16
conv2Conv(16→32) + ReLU + Conv(32→32) + ReLU128×128×32
conv2MaxPool2d(2)64×64×32
conv3Conv(32→128) + ReLU64×64×128
展平view(batch, -1)128×64×64 = 524288 维向量
outLinear(524288 → 20)20(每个类别一个分数)

三个关键组件:

  1. 卷积层(Conv2d):用卷积核在图片上滑动,提取局部特征。浅层提取边缘、纹理等低级特征,深层提取更抽象的语义特征;
  2. 激活函数 ReLUReLU(x) = max(0, x),给网络引入非线性。如果没有激活函数,多层线性变换叠加还是线性,无法拟合复杂函数;
  3. 最大池化 MaxPool2d(2):在每个 2×2 窗口内取最大值。作用是下采样(尺寸减半)、降低计算量、扩大感受野,并让模型对小的位置偏移更鲁棒。

全连接层(Linear):把展平后的 524288 维特征映射到 20 个类别分数(logits)。x.view(x.size(0), -1) 的作用是把 [batch, 128, 64, 64] 展平成 [batch, 524288]-1 表示由 PyTorch 自动推算该维度。


八、训练与评估

8.1 设备选择

device = 'cuda' if torch.cuda.is_available() else 'mps' if torch.backends.mps.is_available() else 'cpu'
model = CNN().to(device)
  • 依次判断是否可用:NVIDIA GPU(cuda)→ Apple 芯片(mps)→ CPU;
  • .to(device) 把模型参数搬到对应设备;训练时数据也要 X.to(device)模型和数据必须在同一设备上

8.2 训练循环:前向 + 反向传播 + 更新

def train(dataloader, model, loss_fn, optimizer):
    model.train()                          # 设为训练模式
    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()                   # ⑤ 按梯度更新参数
        print(f'loss: {loss.item():>7f}')  # 打印当前损失

训练五步曲(每步都很重要):

步骤代码作用
前向传播pred = model(X)图片经过网络得到预测分数
计算损失loss = loss_fn(pred, y)衡量预测与真实标签的差距
梯度清零optimizer.zero_grad()清空上一步的梯度,否则梯度会累加
反向传播loss.backward()用链式法则计算每个参数的梯度
更新参数optimizer.step()优化器按梯度更新模型权重

8.3 损失函数与优化器

loss_fn = nn.CrossEntropyLoss()                        # 多分类交叉熵损失
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)  # Adam 优化器
  • CrossEntropyLoss(交叉熵损失):多分类任务的标准损失函数。它内部已经包含了 Softmax + Log + NLLLoss,所以模型直接输出原始分数(logits)即可,不需要手动加 Softmax;标签必须是 int64 类型的类别编号;
  • Adam 优化器:最常用的自适应学习率优化器,lr=0.001 是它的经典默认学习率。相比普通 SGD,Adam 收敛更快、对学习率不敏感,非常适合入门使用。

8.4 测试/评估

def test(dataloader, model, loss_fn):
    size = len(dataloader.dataset)      # 测试集样本总数
    num_batches = len(dataloader)       # batch 数量
    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(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"Accuracy: {100 * correct}%, Avg loss: {test_loss}")

三个容易忽略的关键点:

  1. model.train() vs model.eval():这两个模式影响 BN(批归一化)和 Dropout 的行为。训练模式用当前 batch 统计量、启用 Dropout;评估模式用全局统计量、关闭 Dropout。评估前必须调用 model.eval()
  2. torch.no_grad():评估时不需要反向传播,用 no_grad 关闭自动求梯度,可以显著节省显存、加快速度;
  3. 准确率计算pred.argmax(1) 取每个样本预测分数最大的类别下标;(pred.argmax(1) == y) 得到布尔张量,.type(torch.float).sum() 转浮点后求和就是预测正确的样本数,除以总数得到准确率。

8.5 完整训练流程

loss_fn = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)

# 先跑一轮,快速查看损失量级
train(train_dataloader, model, loss_fn, optimizer)

# 正式训练 10 个 epoch
epochs = 10
for t in range(epochs):
    print(f"Epoch {t + 1}\n-------------------------------")
    train(train_dataloader, model, loss_fn, optimizer)

# 训练结束后在测试集评估
test(test_dataloader, model, loss_fn)
  • epoch:完整遍历一遍训练集叫 1 个 epoch;
  • 训练 10 个 epoch 只是入门级设置,实际工程中通常观察验证集准确率不再提升(早停)或结合学习率调整来决定训练轮数。

九、保存最优模型

训练结束后,要把模型保存下来,否则关掉程序一切归零。项目里实现了「只保存历史准确率最高的模型」的逻辑:

best_acc = 0    # 记录历史最优准确率

def test(dataloader, model, loss_fn):
    global best_acc
    # ...(计算 test_loss 和 correct,同上一节)...

    if correct > best_acc:               # 当前准确率刷新纪录
        best_acc = correct
        # 方式一:只保存参数
        torch.save(model.state_dict(), 'best_model.pth')
        # 方式二:脚本化保存(结构 + 权重)
        script_model = torch.jit.script(model)
        torch.jit.save(script_model, 'best_model_all.pth')

两种保存方式对比:

对比项torch.save(state_dict)torch.jit.script(model)
保存内容仅模型参数(权重)网络结构 + 参数,打包成一个文件
加载方式需先手动实例化网络类,再 load_state_dict直接 torch.jit.load 使用
是否需要原网络类代码需要不需要
典型用途常规训练/继续训练部署、跨环境使用(TorchScript)

对应的加载代码:

# 方式一:先建模型,再加载参数
model = CNN().to(device)
model.load_state_dict(torch.load('best_model.pth', map_location=device))
model.eval()

# 方式二:直接加载脚本化模型
script_model = torch.jit.load('best_model_all.pth', map_location=device)
script_model.eval()

map_location=device 的作用:把原本保存在 GPU 上的权重映射到当前设备(比如没有 GPU 的机器加载到 CPU),避免设备不匹配报错。


十、单张图片预测

训练好模型后,用一张新图片做预测,这是模型真正"上战场"的环节:

def predict_image(img_path, model, transform, device):
    model.eval()                          # 评估模式
    image = Image.open(img_path).convert("RGB")  # 强制转 RGB,防止灰度图通道数报错
    image = transform(image)              # 预处理:必须与训练/验证时一致
    image = image.unsqueeze(0)            # [C,H,W] → [1,C,H,W],增加 batch 维
    image = image.to(device)
    with torch.no_grad():
        output = model(image)
        pred_class = torch.argmax(output, dim=1).item()  # 取分数最高的类别
    return pred_class

food_names = {
    0: "八宝粥", 1: "哈密瓜", 2: "圣女果", 3: "巴旦木", 4: "板栗",
    5: "汉堡",   6: "火龙果", 7: "炸鸡",   8: "瓜子",   9: "生肉",
    10: "白萝卜", 11: "胡萝卜", 12: "草莓", 13: "菠萝", 14: "薯条",
    15: "蛋",    16: "蛋挞",  17: "菠菜",  18: "骨肉相连", 19: "鸡翅",
}

pred_label = predict_image('test.png', model, data_transforms['valid'], device)
print(f"预测结果:{food_names[pred_label]}")

三个关键细节:

  1. .convert("RGB"):把图片强制转成 RGB 三通道。如果输入是灰度图或带透明通道(RGBA)的 PNG,直接 Image.open 可能得到 1 通道或 4 通道,和模型输入的 3 通道不一致而报错;
  2. .unsqueeze(0):模型输入要求 4 维 [batch, C, H, W]。单张图预处理后是 3 维 [C, H, W],在第 0 维插入一个维度变成 [1, C, H, W]
  3. 预处理必须和训练一致:训练时用哪个 transform 流水线,预测时就要用同一套(尤其归一化参数),否则数据分布不一致会导致预测结果不可信。

十一、完整可运行代码

把上面所有模块整合在一起:

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

# ---------- 1. 数据预处理与增强 ----------
data_transforms = {
    'train': transforms.Compose([
        transforms.Resize([256, 256]),
        transforms.RandomRotation(45),
        transforms.RandomHorizontalFlip(p=0.5),
        transforms.ColorJitter(brightness=0.2, contrast=0.1, saturation=0.1, hue=0.1),
        transforms.ToTensor(),
    ]),
    'valid': transforms.Compose([
        transforms.Resize([256, 256]),
        transforms.ToTensor(),
    ]),
}

# ---------- 2. 自定义数据集 ----------
class food_dataset(Dataset):
    def __init__(self, file_path, transform=None):
        self.imgs = []
        self.labels = []
        self.transform = transform
        with open(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)

    def __len__(self):
        return len(self.imgs)

    def __getitem__(self, idx):
        image = Image.open(self.imgs[idx])
        if self.transform:
            image = self.transform(image)
        label = torch.from_numpy(np.array(self.labels[idx], dtype=np.int64))
        return image, label

# ---------- 3. DataLoader ----------
training_data = food_dataset('train.txt', data_transforms['train'])
test_data = food_dataset('test.txt', data_transforms['valid'])
train_dataloader = DataLoader(training_data, batch_size=64, shuffle=True)
test_dataloader = DataLoader(test_data, batch_size=64, shuffle=True)

# ---------- 4. CNN 模型 ----------
class CNN(nn.Module):
    def __init__(self):
        super(CNN, self).__init__()
        self.conv1 = nn.Sequential(
            nn.Conv2d(3, 16, 5, 1, 2), nn.ReLU(), nn.MaxPool2d(2))
        self.conv2 = nn.Sequential(
            nn.Conv2d(16, 32, 5, 1, 2), nn.ReLU(),
            nn.Conv2d(32, 32, 5, 1, 2), nn.ReLU(),
            nn.MaxPool2d(2))
        self.conv3 = nn.Sequential(
            nn.Conv2d(32, 128, 5, 1, 2), nn.ReLU())
        self.out = nn.Linear(128 * 64 * 64, 20)

    def forward(self, x):
        x = self.conv1(x)
        x = self.conv2(x)
        x = self.conv3(x)
        x = x.view(x.size(0), -1)
        return self.out(x)

device = 'cuda' if torch.cuda.is_available() else 'mps' if torch.backends.mps.is_available() else 'cpu'
model = CNN().to(device)

# ---------- 5. 训练与评估函数 ----------
def train(dataloader, model, loss_fn, optimizer):
    model.train()
    for X, y in dataloader:
        X, y = X.to(device), y.to(device)
        pred = model.forward(X)
        loss = loss_fn(pred, y)
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()
        loss_value = loss.item()
        if batch_size_num % 10 == 0:
            print(f'loss: {loss_value:>7f}  [number:{batch_size_num}]')
        batch_size_num += 1

def test(dataloader, model, loss_fn):
    size = len(dataloader.dataset)
    num_batches = len(dataloader)
    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(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}")
    return correct

loss_fn = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)

# ---------- 6. 训练并保存最优模型 ----------
best_acc = 0
epochs = 10
for t in range(epochs):
    print(f"Epoch {t + 1}\n-------------------------------")
    train(train_dataloader, model, loss_fn, optimizer)
    acc = test(test_dataloader, model, loss_fn)
    if acc > best_acc:
        best_acc = acc
        torch.save(model.state_dict(), 'best_model.pth')
        print(f'已保存新最优模型,准确率: {best_acc * 100:.2f}%')
print("Done!")

# ---------- 7. 单张图片预测 ----------
def predict_image(img_path, model, transform, device):
    model.eval()
    image = Image.open(img_path).convert('RGB')
    image = transform(image).unsqueeze(0).to(device)
    with torch.no_grad():
        pred_class = torch.argmax(model(image), dim=1).item()
    return pred_class

food_names = {
    0: "八宝粥", 1: "哈密瓜", 2: "圣女果", 3: "巴旦木", 4: "板栗",
    5: "汉堡",   6: "火龙果", 7: "炸鸡",   8: "瓜子",   9: "生肉",
    10: "白萝卜", 11: "胡萝卜", 12: "草莓", 13: "菠萝", 14: "薯条",
    15: "蛋",    16: "蛋挞",  17: "菠菜",  18: "骨肉相连", 19: "鸡翅",
}

pred_label = predict_image('test.png', model, data_transforms['valid'], device)
print(f"预测结果:{food_names[pred_label]}")

十二、总结

本项目基于 PyTorch 搭建了一个小型 CNN,完成了 20 类食物的图像分类任务,覆盖了从数据准备到模型预测的完整流程:

  • 数据层:用 os.walk 生成 txt 标签文件,自定义 Dataset(实现 lengetitem)配合 DataLoader 实现批量加载;

  • 预处理层:Resize + ToTensor + Normalize 统一数据格式,训练集额外加入随机旋转、翻转、色彩抖动等数据增强来缓解过拟合;

  • 模型层:3 个卷积块(卷积 + ReLU + 池化)提取特征,全连接层输出 20 类分数;

  • 训练层:CrossEntropyLoss + Adam,遵循"前向 → 损失 → 清零 → 反向 → 更新"五步循环,训练中保存历史最优模型;

  • 应用层:加载模型对单张图片预测,输出对应食物类别。

评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值