PyTorch 食物图像分类实战:从数据准备到 CNN 训练全流程解析

1. 引言

本文围绕一段完整的 PyTorch 食物图像分类代码展开,逐段讲解从数据集整理、自定义 Dataset、数据预处理、CNN 网络搭建到训练与测试循环的完整流程。代码以食物图片二分类(train/test 目录)为背景,适合初学者理解 PyTorch 图像分类任务的基本套路。

2. 环境与依赖

代码依赖以下核心库:

  • os:文件路径遍历与拼接。
  • torch / torch.nn:张量计算与神经网络模块。
  • numpy:数值计算(本代码中主要用于间接支持张量操作)。
  • PIL(Pillow):读取图片文件。
  • torch.utils.data:Dataset 与 DataLoader 数据加载工具。
  • torchvision.transforms:图像预处理与张量转换。

3. 数据集目录整理:train_test_file 函数

该函数的作用是把指定目录下的图片路径和类别标签写入 txt 文件,供后续 Dataset 读取。核心逻辑如下:

python 复制代码
def train_test_file(root, dir):
    # 使用with自动关闭文件,避免文件泄露
    with open(dir + '.txt', 'w') as file_txt:
        path = os.path.join(root, dir)
        dirs_cache = []  # 缓存类别文件夹,修复NameError问题
        for roots, directories, files in os.walk(path):
            if len(directories) != 0:
                dirs_cache = directories
            else:
                now_dir = roots.split(os.sep)  # os.sep自动适配Windows\\和Linux/
                for file in files:
                    path_1 = os.path.join(roots, file)
                    print(path_1)
                    # 防止类别不存在报错
                    if now_dir[-1] in dirs_cache:
                        file_txt.write(path_1 + ' ' + str(dirs_cache.index(now_dir[-1])) + '\n')

逐行解析:

  • with open(dir + '.txt', 'w'):以写模式打开文件,with 语句确保文件自动关闭,避免资源泄露。
  • os.walk(path):递归遍历目录,返回三元组(当前目录路径、子目录列表、文件列表)。
  • dirs_cache = directories:当当前目录还有子目录时,缓存类别文件夹名。这里修复了原代码中可能出现的 NameError 问题------当进入叶子目录时,directories 为空,需要用之前缓存的类别列表。
  • roots.split(os.sep):按系统分隔符拆分路径,os.sep 在 Windows 下是反斜杠,在 Linux 下是正斜杠,保证跨平台兼容。
  • now_dir-1 in dirs_cache:判断当前叶子目录名是否属于已知类别,防止类别不存在时报错。
  • dirs_cache.index(now_dir-1):取类别在列表中的下标作为数字标签,写入 txt 文件。

最终生成的 train.txt / test.txt 每行格式为:图片绝对路径 + 空格 + 类别编号。

4. 自定义 Dataset 与 getitem 机制

代码先通过一个简单的 USE_getitem 类演示了 Python 的索引协议:

python 复制代码
class USE_getitem():
    def __init__(self, text):
        self.text = text
    def __getitem__(self, index):
        result = self.text[index].upper()
        return result
    def __len__(self):
        return len(self.text)

这个类实现了两个魔法方法:getitem 让对象支持下标访问(p1),len 让对象支持 len() 调用。这是 PyTorch Dataset 的核心协议------DataLoader 正是通过这两个方法按索引取样本并统计长度。

接下来是真正的食物数据集类:

python 复制代码
class food_dataset(Dataset):
    def __init__(self, file_path, transform=None):
        self.file_path = file_path
        self.imgs = []
        self.labels = []
        self.transform = transform
        with open(self.file_path, 'r', encoding='gbk') 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(int(label))  # 直接转int,不再存字符串
    def __len__(self):
        return len(self.imgs)
    def __getitem__(self, idx):
        image = Image.open(self.imgs[idx]).convert("RGB")  # 统一转RGB,解决灰度图通道报错
        if self.transform:
            image = self.transform(image)
        label = self.labels[idx]
        return image, label

关键点:

  • init:读取 txt 文件,按空格拆分每行,分别存入图片路径列表和标签列表。标签直接转为 int,避免后续计算类型不匹配。
  • getitem:用 PIL 打开图片并统一转为 RGB 三通道,解决灰度图导致的通道数报错问题;随后应用 transform 预处理,返回(图像张量,标签)元组。
  • len:返回样本总数,供 DataLoader 计算迭代轮次。

5. 图像预处理与 DataLoader

python 复制代码
data_transforms = {
    # 机器学习的时候,数据进行归一化? 图片做归一化 ->0~1
    'trainda':
        transforms.Compose([
            transforms.Resize([256, 256]),  # 图像变换大小 opencv int8 0~255
            transforms.RandomRotation(45),  # 随机旋转,-45到45度之间随机选
            transforms.CenterCrop(256),  # 从中心开始裁剪[256,256]
            transforms.RandomHorizontalFlip(p=0.5),  # 随机水平翻转 选择一个概率
            transforms.RandomVerticalFlip(p=0.5),  # 随机垂直翻转
            transforms.ColorJitter(brightness=0.2, contrast=0.1, saturation=0.1, hue=0.1),  # 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相
            transforms.RandomGrayscale(p=0.1),  # 概率转换成灰度,3通道就是R=G=B
            transforms.ToTensor(),
            transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])  # 标准化,均值,标准差
        ]),
    'valid':  # 验证集 不需要对图像进行数据增强
        transforms.Compose([
            transforms.Resize([256, 256]),
            transforms.ToTensor(),
            transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
        ]),
}

transforms.Compose 把多个预处理操作串联成一个管道:

  • Resize(256, 256):把图片统一缩放到 256×256,保证网络输入尺寸一致。
  • RandomRotation(45):在 -45 到 45 度之间随机旋转图片,增强模型对旋转变化的鲁棒性。
  • CenterCrop(256):从中心裁剪出 256×256 的区域,配合 Resize 保证输入尺寸一致。
  • RandomHorizontalFlip(p=0.5):以 0.5 的概率随机水平翻转图片。
  • RandomVerticalFlip(p=0.5):以 0.5 的概率随机垂直翻转图片。
  • ColorJitter(brightness=0.2, contrast=0.1, saturation=0.1, hue=0.1):随机调整亮度、对比度、饱和度和色相,增强模型对颜色变化的鲁棒性。
  • RandomGrayscale(p=0.1):以 0.1 的概率将图片转为灰度图(三通道 R=G=B)。
  • ToTensor():把 PIL 图像或 numpy 数组转为 PyTorch 张量,并将像素值从 0-255 归一化到 0-1,同时把通道维度从 HWC 调整为 CHW。
  • Normalize(mean, std):用 ImageNet 的均值和标准差对张量做标准化,加速模型收敛。

训练集使用完整的数据增强管道,验证集只做 Resize、ToTensor 和 Normalize,不做随机增强,保证评估结果稳定。

随后创建数据集和加载器:

python 复制代码
training_data = food_dataset('./train.txt', transform=data_transforms['trainda'])
test_data = food_dataset('./test.txt', transform=data_transforms['valid'])
train_dataloader = DataLoader(training_data, batch_size=4, shuffle=True)
test_dataloader = DataLoader(test_data, batch_size=4, shuffle=True)

DataLoader 负责把 Dataset 包装成可迭代对象:batch_size=4 表示每次取 4 个样本组成一个 batch;shuffle=True 表示每个 epoch 开始前打乱数据顺序,提升训练稳定性。

6. CNN 网络结构解析

python 复制代码
class CNN(nn.Module):
    def __init__(self):
        super(CNN, self).__init__()
        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),
        )
        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)
        output = self.out(x)
        return output

网络由三个卷积块和一个全连接输出层组成,逐层分析:

  • conv1:输入 3 通道(RGB),输出 16 通道,卷积核 5×5,padding=2 保持尺寸不变,ReLU 激活后接 2×2 最大池化,特征图从 256×256 降为 128×128。
  • conv2:16→32 通道卷积 + ReLU,再接 32→32 卷积 + ReLU,最后 2×2 池化,特征图降为 64×64。
  • conv3:32→128 通道卷积 + ReLU,无池化,特征图保持 64×64。
  • out:全连接层,输入维度为 128×64×64(展平后的特征总数),输出 20 个类别分数。

forward 方法中,x.view(x.size(0), -1) 把每个样本的 128×64×64 特征图展平为一维向量,再送入全连接层得到 20 类别的 logits。

7. 训练与测试循环

python 复制代码
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = CNN().to(device)
loss_fn = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)

这段代码完成三件事:

  • device 选择:优先使用 GPU(cuda),否则回退到 CPU。
  • 模型与损失函数:实例化 CNN 并移动到指定设备;CrossEntropyLoss 内部包含 Softmax 和交叉熵计算,适合多分类任务。
  • 优化器:Adam 优化器,学习率 1e-4,更新模型全部参数。
python 复制代码
def train_loop(dataloader, model, loss_fn, optimizer):
    model.train()
    for batch, (X, y) in enumerate(dataloader):
        X, y = X.to(device), y.to(device)
        pred = model(X)
        loss = loss_fn(pred, y)
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()
        if batch % 10 == 0:
            print(f"Train Batch:{batch}, Loss: {loss.item():.4f}")

训练循环的标准五步:

  • model.train():切换到训练模式,启用 Dropout 和 BatchNorm 的训练行为。
  • 前向传播:把 batch 数据送入模型得到预测 pred。
  • 计算损失:pred 与真实标签 y 计算交叉熵。
  • 反向传播:optimizer.zero_grad() 清空旧梯度,loss.backward() 计算新梯度,optimizer.step() 更新参数。
  • 日志输出:每 10 个 batch 打印一次当前损失,便于观察收敛情况。
python 复制代码
def test_loop(dataloader, model, loss_fn):
    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 /= len(dataloader)
    correct /= len(dataloader.dataset)
    print(f"==== Test ==== Loss: {test_loss:.4f}, Acc: {correct:.4f}")

测试循环与训练的关键区别:

  • model.eval():切换到评估模式,关闭 Dropout 等训练专用行为。
  • torch.no_grad():关闭梯度计算,节省内存并加速推理。
  • pred.argmax(1):取每个样本预测分数最高的类别下标,与真实标签比较,统计正确数。
  • 准确率计算:正确数除以总样本数,得到测试集准确率。

8. 完整训练流程与注意事项

把上述函数串联起来即可开始训练:

python 复制代码
epochs = 10
for t in range(epochs):
    print(f"Epoch {t+1}\n-------------------------------")
    train_loop(train_dataloader, model, loss_fn, optimizer)
    test_loop(test_dataloader, model, loss_fn)
print("Done!")

实际使用中需要注意以下几点:

  • 类别数匹配:全连接层输出 20 对应 20 个类别,需与数据集实际类别数一致。
  • 内存占用:128×64×64 的展平特征较大,若显存不足可减小 batch_size 或增加池化层。
  • 数据增强:当前仅做了 Resize 和 ToTensor,可加入 RandomHorizontalFlip、RandomRotation 等增强手段提升泛化能力。
  • 学习率调整:1e-4 是较保守的初始值,训练后期可配合学习率衰减策略。

9. 总结

本文完整解析了 PyTorch 食物图像分类的代码链路:先用 os.walk 整理数据集生成标签文件,再通过自定义 Dataset 读取图片和标签,配合 transforms 预处理和 DataLoader 批量加载,最后搭建三层 CNN 并完成训练与测试循环。理解这套流程后,可以轻松迁移到其他图像分类任务,只需调整类别数、网络结构和数据增强策略即可。

相关推荐
词却1 小时前
深度学习入门:卷积神经网络与 MNIST 手写数字识别
人工智能·深度学习·cnn
adaierya2 小时前
用 AI 解决音频转换编程问题
开发语言·人工智能·python·分类·ai编程
ai小陈2 小时前
PyTorch梯度累积与裁剪实战:小显存也能稳定训练大批次
人工智能·pytorch·python·深度学习·ai·gpu算力
C^h11 小时前
pytorch 适合初学者 0基础学习
人工智能·pytorch·python
论文复现现场20 小时前
本地 PyTorch 训练 OOM,第一次租 RTX 4090 云 GPU 怎么迁移项目?从环境检查到 100 Step 跑通
人工智能·pytorch·python·深度学习·cuda
2601_9620990821 小时前
在Python中使用LSTM和PyTorch进行时间序列预测
pytorch·机器学习·lstm·时间序列预测·数据预处理
论文复现现场21 小时前
单卡RTX 3090能训练,切到4卡却OOM:Accelerate多卡训练怎么排查?
人工智能·pytorch·云计算·gpu算力·多卡训练
小白说大模型1 天前
去AI味提示词大全:25个实用Prompt帮你降低AI率
大数据·人工智能·pytorch·深度学习·机器学习·prompt
Hali_Botebie1 天前
PyTorch 内存布局,.view()要合并哪两个维度(比如 B 和 G),这两个维度在内存里就必须“紧挨着”。
人工智能·pytorch·python