☰
PyTorch猫狗公鸡三分类实战:环境配置、数据增强与训练排错指南
2026/9/28 2:02:06 网站建设 项目流程

简介:为希望掌握PyTorch图像分类流程的开发者提供一份完整实践项目,以猫、狗、公鸡三分类为具体案例,系统讲解从数据预处理、CNN网络搭建到训练验证的完整过程,适合有一定深度学习基础、想通过实战系统上手PyTorch的初学者。压缩包内共1390个文件,其中1362张jpg构成训练与验证图片数据集,11个py脚本是核心训练与推理代码,另含5个xml配置文件、3个txt说明文档、1个onnx模型文件等,包体554.92MB,目录结构清晰,便于按步骤学习。已有1378人参与学习下载,项目内容广受认可。通过这套实战项目,可以动手实践数据读取与增强、Conv2d卷积层堆叠、CrossEntropyLoss损失函数与SGD/Adam优化器选择、DataLoader数据加载、模型保存与加载等关键操作,并可借助TensorBoard或混淆矩阵分析分类效果,从而快速掌握PyTorch在图像分类任务中的典型开发流程,为后续更复杂的深度学习项目打下坚实基础。

1. 用 PyTorch 做猫狗公鸡图片分类:三分类比二分类更值得练手

把猫狗分类换成猫狗公鸡三分类,难度马上不一样。狗和公鸡的毛色、纹理与姿态有大量重叠,网络稍微偷懒就会去记背景而不是动物本体。用 PyTorch 搭建一个猫狗公鸡图片分类网络,核心不是套一个现成模型跑一遍,而是把目录结构、数据增强、网络设计、训练循环和排错手段串成一条能复现的链路。这篇笔记适合刚学完 Python 与 PyTorch 基础、准备做第一个实战项目的人,也适合已经跑过二分类但想给自己加一点难度的从业者。我们按“先跑通环境—准备数据—搭网络—训练—排坑—验证”的顺序走完,每条命令和参数都能直接抄。

2. 先跑通环境再谈模型:PyTorch 安装与开发环境配置

新手一开始就把精力花在挑网络结构上,这是顺序问题。环境搭不稳,后面每一个报错都会混在一起:分不清是代码错了、包版本错了,还是 GPU 没调用起来。我一般把环境配置控制在二十分钟内解决,目标只有一个:能跑通一个最简单的张量运算,确认 torch 和 torchvision 都能正常 import。

2.1 用 Anaconda 建独立环境:避免依赖打架

很多人第一次装 PyTorch 是在系统 Python 里直接 pip install torch,装到后面发现图片分类代码一运行就报错。最常见的原因不是代码写错,而是 torch 与 torchvision 的版本不匹配,或者和系统里其他包冲突。PyTorch 的官方约束很明确:torchvision 是配套 torch 单独发布的,版本必须对应,不能单独升级其中一个。

推荐先用 Anaconda 建一个干净环境,把这个项目的依赖和系统其他 Python 环境隔开。这个习惯在多个项目并行时能省下大量排错时间:项目 A 升级了 torch,不会影响项目 B 的依赖树。

conda create -n catdog python=3.10 -y conda activate catdog

python=3.10 是目前兼容性比较稳的选择,支持 torch 的同时也能满足大部分图像处理库的版本要求。如果你本机已经跑着 3.11 或 3.12 的现成项目,不冲突时也可以用,但建议按这个命令新起一个环境,避免以后换显卡驱动或装新库时把 base 环境搞坏。如果你不想装 Anaconda,用 Python 自带的 venv 也能做隔离,只是后续切换环境和安装 CUDA 相关依赖时,conda 处理起来更省心。

为什么不直接 conda install torch?常见原因是 conda 默认源在一些机器上很慢,而且 conda 对 torch 这种大型二进制包的依赖解析比 pip 慢不少。pip 安装时 torch 与 torchvision 会校验彼此的版本约束,装错会即时报错,比跑到训练时才暴露问题好处理。

2.2 CPU 版还是 GPU 版:三个判断条件

很多人上来就问“我该装 CPU 版还是 GPU 版”,这个问题在猫狗公鸡三分类场景里没那么复杂。几十到几百张的小样本量,CPU 版完全可以跑完整个训练流程,只是慢一点。判断要不要花时间装 GPU 版,只看三件事:有没有 NVIDIA 显卡;显卡驱动支持的 CUDA 版本能不能和 PyTorch 对上;你后面是不是要反复训练调参。三条中有一条不确定,就先装 CPU 版把逻辑跑通,再换 GPU 版,不要一上来就追 CUDA 配置。

常见做法是 CPU 版直接 pip 安装:

pip install torch torchvision

如果需要 GPU 版,安装命令会指定一个带 CUDA 版本的下载源:

pip install torch torchvision --index-url https://download.pytorch.org/whl/cu121

cu121 对应 CUDA 12.1 的预编译版本,PyTorch 官方安装页会给不同 CUDA 版本列出对应命令。新手不理解的坑是:这里选哪个 CUDA 版本不是看 PyTorch 版本号,而是看显卡驱动支持到什么级别。装完如果 import 正常但 GPU 不可用,绝大多数情况是驱动版本高于或低于 PyTorch 打包的 CUDA 运行库,先试官方推荐的默认版本,不行再换低一档的 cu118。老显卡不确定时,敲一下nvidia-smi,看右上角显示的 CUDA 版本,拿它减去 0.5 再向下取整,基本就是能用的档位。

2.3 安装完必须做的三个验证

环境装完先别急着写网络,用一段最简单的代码验证三件事:PyTorch 版本能读到、CUDA 可用性返回值符合预期、torchvision 能正常导入。这个验证以后每次换机器都能用。

import torch print(torch.__version__) print(torch.cuda.is_available()) from torchvision import transforms, datasets print("torchvision ok")

如果 torch.cuda.is_available() 返回 False,不一定是装错了。先看 pip list 里 torch 版本号是否带 +cpu 后缀,带就是 CPU 版,不带再看驱动。很多人在这一步死磕 GPU,结果发现自己的显卡其实不支持当前 CUDA 版本,白折腾一小时。小样本三分类用 CPU 训练二十轮也就几分钟,先把网络跑起来,再回头优化算力成本,性价比高得多。

另外,如果你在 WSL 里搭环境,验证命令和普通 Linux/Windows 完全一样,PyTorch 对 WSL 的支持已经很成熟。只有一个细节:WSL 里的 conda 环境和 Windows 侧不共享,别两边反复装。torch 和 torchvision 的版本对应关系可以用pip list | grep torch确认,主版本号一致基本就没问题。

装完这两个核心包后,这个项目基本不需要再引入其他深度学习库。数据操作用 torchvision,可视化用 matplotlib,如果只是写脚本训练,这两步就够了。如果想确认网络每层的输出尺寸,也不用额外装包,直接print(model)就能看到每一层的名字和参数形状,新手阶段这个输出比任何工具都直观。

3. 猫狗公鸡数据准备:目录结构、图像预处理与数据增强

数据准备是图片分类里最容易被低估的一步。很多新手把图片随便丢到一个文件夹里,然后开始写模型,最后训练时要么报“Found 0 images”,要么模型在验证集上过拟合得一塌糊涂。数据这一步决定了后面所有环节的上限,网络设计解决不了数据没整理好的问题。

3.1 目录结构:用 ImageFolder 按文件夹生成标签

torchvision 里的datasets.ImageFolder就是为这类目录结构设计的。它把每个子目录当作一个类别,根据图片文件自动生成标签,不需要你手动维护一份标签表。对猫狗公鸡三分类来说,这是最快、最不容易错的方案。

data/ ├── train/ │ ├── cat/ # 猫图片 │ ├── dog/ # 狗图片 │ └── rooster/ # 公鸡图片 └── val/ ├── cat/ ├── dog/ └── rooster/

train 和 val 必须分开,而且 val 里不能混入 train 的图片,否则验证结果会虚高,后面说避坑时还会提到。ImageFolder 读取时会按字母顺序给子目录排序,cat 对应索引 0、dog 对应 1、rooster 对应 2。这个顺序在推理阶段很重要:保存模型时把 classes 输出记下来,以后推理类别索引以它为准。

from torchvision import datasets train_dataset = datasets.ImageFolder(root='data/train') print(train_dataset.classes) print(train_dataset.class_to_idx)

如果发现 class_to_idx 的顺序和自己预期不一致,不要改文件夹名去迎合,因为排序规则是固定的,改名字反而容易在换数据时翻车。直接用 classes 列表的索引解读预测结果即可。

3.2 预处理管线:Resize、ToTensor 与 Normalize 的顺序

图片读进来是 PIL 对象,取值范围 0 到 255,形状是 HWC。PyTorch 网络要求输入是 CHW 顺序的浮点张量,像素值最好落在 0 到 1 附近,所以 transforms 的顺序有严格规定:先做几何变换,再 ToTensor,最后 Normalize。

from torchvision import transforms transform_train = transforms.Compose([ transforms.RandomResizedCrop(224, scale=(0.5, 1.0)), transforms.RandomHorizontalFlip(p=0.5), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) transform_val = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ])

train 用了 RandomResizedCrop,会在原图上随机裁剪一块区域再缩放成 224x224,相当于同时完成了裁剪和缩放;val 用 Resize(256) 再 CenterCrop(224),保证验证时每张图都用固定的中心区域,结果可复现。mean 和 std 沿用 ImageNet 的统计值,这是 torchvision 预训练模型的标准预处理,自己从零训练时沿用也不会错。224 是 torchvision 预训练模型默认的输入尺寸,也是显存和精度比较平衡的选择。

如果你自己收集的图片尺寸差异很大,Resize 时要注意:直接用 Resize((224,224)) 会把长宽比压扁,公鸡被拉宽后纹理形变会影响分类。训练阶段用 RandomResizedCrop 天然解决了比例问题,验证阶段用先 Resize 短边再 CenterCrop 的固定流程,比直接强压更稳。

3.3 数据增强:小数据集的三件套与边界

猫狗公鸡三分类如果每个类别只有几十张图,不加增强基本必过拟合。增强的本质是给模型制造合法扰动:公鸡换个角度、光线变一点、颜色偏移一点,分类结果不能变。但增强不是越多越好,旋转超过 30 度连人眼都难分辨,网络学到的就变成了“猜”。

transform_train = transforms.Compose([ transforms.RandomResizedCrop(224, scale=(0.5, 1.0)), transforms.RandomHorizontalFlip(p=0.5), transforms.RandomRotation(15), transforms.ColorJitter(brightness=0.2, contrast=0.2), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ])

RandomRotation(15) 的 15 表示角度范围正负 15 度。ColorJitter 里 brightness、contrast 各 0.2,表示在 0.8 到 1.2 倍之间随机调整。对公鸡这类纹理敏感的类别,旋转和颜色扰动比翻转更重要,因为公鸡的鸡冠和尾羽在不同角度下变化很大,网络容易把这些外观细节当成固定特征记死。这套增强对几十张的小数据集足够,不用再加更复杂的 AutoAugment。

如果公鸡类别实在只有三四十张图,我会在 Normalize 之后再加一层 RandomErasing。它会在输入图上随机抹掉一块矩形区域,让网络不要过度依赖某一个局部特征,比如鸡冠。p 参数控制在 0.3 以内,scale 设为 (0.02, 0.1),意思是抹除面积占原图 2% 到 10%。这块是血泪经验:小数据集上不加遮挡类增强,验证集只要带一点遮挡,准确率就崩。

3.4 DataLoader:batch_size、num_workers 与样本不均衡

数据准备好之后用 DataLoader 包起来,训练循环每次从里面取一批图片。参数不多,但每一个都直接影响训练速度和显存占用。

from torch.utils.data import DataLoader train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=4, pin_memory=True) val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=4, pin_memory=True)

batch_size=32 表示每批 32 张图算一次梯度。显存有限时第一个调整的就是它:降到 16 或 8,训练变慢但不会爆显存。shuffle=True 只在训练集开,验证集保持固定顺序,否则每个 epoch 的验证样本顺序一直在变,不方便对比。num_workers 是读取图片的子进程数,默认 0 表示用主进程读图,数据稍微多点训练就卡在 IO 上;设成 CPU 核数减一比较合理,Windows 上建议不要超过 4,偶发多进程兼容问题能少一点。pin_memory=True 在 GPU 训练时能把数据放到锁页内存,减少 CPU 到 GPU 的拷贝开销,CPU 训练也不会有副作用。

还有一个容易被忽略的问题:如果猫图片有 500 张,公鸡只有 50 张,ImageFolder 会按顺序喂数据,模型会严重偏向猫。常见做法是给 DataLoader 加 WeightedRandomSampler,按类别数量倒数的权重采样,让小类别每个 epoch 也能被抽到足够多次。这个方向知道即可,先把均衡数据跑通,再回头看是否需要加权。

4. 搭建图片分类网络:卷积层、分类头与训练闭环

网络结构是这类项目里最容易被神话的部分。实际上,猫狗公鸡三分类任务用三层卷积已经能到 90% 以上准确率,前提是数据流程正确、训练参数合理。先学会把网络每一层的输入输出算清楚,再谈换 ResNet 是更稳的路线。

4.1 手写一个三层卷积网络:从 Conv2d 到 MaxPool2d

输入是 3 通道的 224x224 图像,输出是 3 个类别的得分向量。中间的处理逻辑是:卷积层在局部窗口里提取纹理特征,ReLU 做非线性变换,池化层把特征图缩小一半,最后用全连接层把特征映射成类别得分。这个结构的每一层尺寸变化是可以手算的,写代码前先在纸上过一遍。

import torch import torch.nn as nn class CatDogRoosterNet(nn.Module): def __init__(self, num_classes=3): super().__init__() self.features = nn.Sequential( nn.Conv2d(3, 16, kernel_size=3, padding=1), nn.ReLU(inplace=True), nn.MaxPool2d(2), nn.Conv2d(16, 32, kernel_size=3, padding=1), nn.ReLU(inplace=True), nn.MaxPool2d(2), nn.Conv2d(32, 64, kernel_size=3, padding=1), nn.ReLU(inplace=True), nn.MaxPool2d(2), ) self.classifier = nn.Sequential( nn.Flatten(), nn.Linear(64 * 28 * 28, 256), nn.ReLU(inplace=True), nn.Linear(256, num_classes), ) def forward(self, x): x = self.features(x) x = self.classifier(x) return x

尺寸变化是理解这段代码的关键。224x224 的图像经过第一层卷积,padding=1 保证输出仍是 224x224,然后 MaxPool2d(2) 后变成 112x112。第二层卷积后池化成 56x56,第三层后是 28x28。通道数从 3 变成 16、32、64,所以 Flatten 后传给全连接层的特征维度是 64 乘 28 乘 28,也就是 50176。如果改了输入尺寸或加了池化层,这个数字必须跟着改,这是新手最常见的报错来源之一。

ReLU 加 inplace=True 是省内存的写法,意思是在原张量上直接修改,不新建输出张量。BatchNorm 的位置有讲究,常见的是 Conv 之后、ReLU 之前放一层 BatchNorm2d。但我们的数据已经做了归一化,小网络里不写 BatchNorm 也能稳定训练,还能少踩一个 eval 模式的坑,所以我在这份代码里先不写 BatchNorm,等你迁移到更深网络时再补。

4.2 分类头设计:为什么输出三个神经元

最后的全连接层输出 3 个数值,对应猫、狗、公鸡三个类别的得分。训练时用 CrossEntropyLoss,它会把这些得分(logits)做 softmax 变成概率,再和真实标签做交叉熵计算。这里有一个容易翻车的点:CrossEntropyLoss 内部已经包含 LogSoftmax,所以不要在最后一层手动加 Softmax 或 Sigmoid,否则梯度路径会出问题,训练半天 loss 不降。

self.classifier = nn.Sequential( nn.Flatten(), nn.Linear(64 * 28 * 28, 256), nn.ReLU(inplace=True), nn.Dropout(0.2), nn.Linear(256, 3), )

Dropout(0.2) 会在训练时随机让 20% 的神经元输出变成 0,强迫网络不依赖单个神经元。它对小数据集防过拟合很有效,但带来的副产品是:验证时必须调用model.eval()关闭 Dropout,否则推理结果是随机的,这一点稍后讲。如果只是做三分类,中间隐藏层 256 已经够用,加到 512 不会带来明显收益,反而更慢。

4.3 迁移学习:用预训练模型替换分类头

自己从零搭网络最大的好处是理解每一层在干什么,但真实项目的生产力方案是迁移学习。torchvision 里带预训练权重,在 ImageNet 上学过的纹理特征可以直接复用,你只需要把最后一层换成输出为 3 的全连接层。

import torchvision.models as models backbone = models.resnet18(weights=models.ResNet18_Weights.DEFAULT) backbone.fc = nn.Linear(backbone.fc.in_features, 3)

关键是以backbone.fc.in_features为准去替换,不要写死 512,不同模型的分类头输入维度不一样。预训练权重要求输入尺寸和归一化方式与 ImageNet 一致,所以前面 transform 里的 Resize 224、mean/std 那套必须保持原样。如果数据集很小,可以先冻结 backbone 的卷积层只训练分类头,代码上就是过滤掉不需要更新参数的层:

for param in backbone.parameters(): param.requires_grad = False for param in backbone.fc.parameters(): param.requires_grad = True

冻结之后再训练,loss 下降会比全网络微调慢,但不容易过拟合。当你试完手写网络、理解了卷积和全连接的配合,再用迁移学习去对比一次精度提升,比直接上 ResNet 的收获大得多。

4.4 训练循环:把 loss、梯度更新和学习率调度串起来

网络写好后,训练循环是固定套路,但代码顺序不能错。完整看一下一个 epoch 的训练代码:

epochs = 10 criterion = nn.CrossEntropyLoss() optimizer = torch.optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-4) scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=3, gamma=0.5) for epoch in range(epochs): model.train() total_loss = 0 for images, labels in train_loader: optimizer.zero_grad() outputs = model(images) loss = criterion(outputs, labels) loss.backward() optimizer.step() total_loss += loss.item() avg_loss = total_loss / len(train_loader) scheduler.step() print(f"epoch {epoch + 1:02d}, train loss {avg_loss:.4f}, " f"lr {scheduler.get_last_lr()[0]:.2e}")

optimizer.zero_grad() 必须写在 loss.backward() 之前,作用是清空上一次迭代留下的梯度。如果漏了,梯度会在每次迭代里累加,模型参数更新方向越来越偏。loss.backward() 计算梯度后,optimizer.step() 更新参数。model.train() 和 model.eval() 的切换也要注意:每轮训练开始切 train,验证前切 eval。

学习率调度选 StepLR 是常见的保守做法:每 3 个 epoch 学习率减半。Adam 的初始学习率 1e-3 对这个小网络是安全起点。如果你想用 SGD,一般从 0.01 起步配 momentum=0.9。判断学习率是否合理,直接观察第一个 epoch 的 loss 变化:如果第一个 epoch 结束 loss 从 1.1 降到 0.7 左右,说明学得动;如果纹丝不动,大概率学习率太小或数据链路有问题,不是模型问题。

训练时还要同时关注验证集。每轮结束后跑一次验证,记录 val loss 和准确率,才能判断模型是正在拟合还是已经过拟合。验证代码用 model.eval() 加 torch.no_grad(),是固定搭配,少一个结果都不可信。

5. 训练踩坑与常见问题排查:5 个让模型翻车的问题

训练跑起来只是开始,真正的工程量在排错。下面这几类问题我在带新手做图片分类时反复见到,每条都按现象、原因、解决来写,你可以直接把日志和情况对号入座。这些问题单独看都不难,但组合出现时特别费时间,所以把边界条件也写清楚:什么情况算正常,什么情况必须停。

5.1 Loss 停在 1.10 附近不动

现象:训练 loss 从初始值缓慢下降后停在 1.09 到 1.12 之间,验证集准确率一直在 33% 左右。1.099 是三分类随机猜测的理论交叉熵值,等于 -ln(1/3)。看到这个数字说明模型没在学东西,而不是“学得慢”。

原因:最常见是学习率实在太小,梯度更新幅度不足以改变权重;其次是数据链路出错,比如所有图片都被放进了同一个类别文件夹,标签和图像对应不上;还有一种常见情况是 Normalize 的 mean 和 std 填反了,输入分布不在模型预期范围,训练直接失效。

解决:先打印一个 batch 的 images.mean() 和 images.std(),确认数值范围正常;再把学习率从 1e-3 上调到 1e-2(Adam)或从 0.01 起调(SGD),各跑 5 个 epoch 对比;最后用torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)限制梯度范数,防止个别样本把梯度拉爆。

5.2 验证准确率比训练还高

现象:训练集准确率在区间内波动,验证集准确率反而稳稳高出十几个点,看起来很反常。

原因:多数情况是验证循环里忘了加 model.eval()。网络里有 Dropout 或 BatchNorm 时,训练模式会用当前 batch 的统计量做归一化;验证时不切 eval,BatchNorm 还在用训练状态计算,输出并不是真正的模型表现。另一个原因确认一下目录:val 集里混进了 train 的图片,这是数据泄露。

解决:验证循环固定写成下面这个模式,两个字都不能少:

model.eval() with torch.no_grad(): # 跑完整验证集,只统计数据,不更新梯度 correct = 0 total = 0 for images, labels in val_loader: outputs = model(images) _, predicted = torch.max(outputs, 1) total += labels.size(0) correct += (predicted == labels).sum().item()

每轮验证结束后,记得把模型切回model.train(),否则下一轮训练的 Dropout 不生效,loss 曲线会突然出现毛刺。

5.3 训练到中途报 CUDA out of memory

现象:第一个 epoch 顺利跑完,第二个 epoch 或验证阶段报RuntimeError: CUDA out of memory。很多人以为是显存被占满,实际是验证阶段把整个 val 集一次性全塞进去,或者优化器的梯度、动量等辅助内存逐步累积。

原因:batch_size 设置过大,224x224 的激活值占用高;验证集没有分 batch 推理;多个进程共享同一块显存但没有释放。

解决:优先把 batch_size 从 32 降到 16 或 8,这一步通常立竿见影;再把输入尺寸统一缩到 160,对三分类精度影响很小;epoch 结束时调用torch.cuda.empty_cache()主动释放缓存。如果还紧张,把模型通道数从 64 降到 32,这类小任务完全够用。

5.4 公鸡总被认成狗

现象:整体准确率到 80% 以上,但混淆矩阵里公鸡类只有 60%,大部分错判成狗。整体 acc 掩盖了单类问题,只看整体指标发现不了。

原因:公鸡的鸡冠、羽毛纹理在部分姿态下和狗的卷毛非常接近;更常见的是公鸡图片来源单一,全都带草地或栅栏背景,模型学到的是“浅绿背景”而不是公鸡本身。测试场景一换,准确率立刻下降。

解决:给 rooster 类增加多背景、多角度的图片,这是最本质的解法;增强里开启 RandomErasing 和 ColorJitter,打掉颜色依赖;如果训练图背景干扰明显,先用检测框把动物主体裁出来再训练。想打印混淆矩阵,用 sklearn.metrics.confusion_matrix,把验证集的预测值和真实值传进去即可,肉眼看一下错在哪一类,比任何调参都有方向。调网络结构解决不了数据分布问题,因为模型记住的是背景里稳定的颜色块,不是纹理。

5.5 ImageFolder 报 “Found 0 images”

现象:DataLoader 初始化直接抛RuntimeError: Found 0 images in subfolders of ...,路径检查了好几遍看不出问题。

原因:路径不存在、子目录层级比预期多一层,或者目录里混入了 .DS_Store、Thumbs.db 这类系统隐藏文件,被 ImageFolder 当成了类别文件夹。分类网络代码本身没问题,是数据目录不干净。

解决:先用os.listdir('data/train')打印顶层目录,确认只有 cat、dog、rooster 三个文件夹;进到每个类别目录里用ls | head -5看图片文件列表,隐藏文件全部清理掉。Windows 上路径尽量别用中文和空格,PyTorch 对 unicode 路径的处理偶尔会有兼容问题。Linux 上注意扩展名大小写敏感,.JPG 和 .jpg 是不同文件,ImageFolder 对 PIL 能打开的后缀都能识别,但目录名的大小写必须和代码一致。

6. 模型保存与推理验证:把训练好的网络真正用起来

训练完这一步几乎所有人都会做:保存模型。但保存方式有讲究。推荐只保存 state_dict,它只含权重字典,体积小且换环境不依赖模型类定义文件。

torch.save(model.state_dict(), "cat_dog_rooster.pth")

加载时先构建网络再填权重,注意网络结构必须和训练时完全一致。模型文件拷到没有 GPU 的机器上时,加 map_location 参数:

model = CatDogRoosterNet(num_classes=3) model.load_state_dict(torch.load("cat_dog_rooster.pth", map_location="cpu")) model.eval()

接着写一个推理函数,输入图片路径,输出各类别概率。class_names 来自训练时的train_dataset.classes。这里有一个我一直会加的技巧:打印 top-2 的概率而不是只打印最大概率。当第一和第二名的概率接近,比如 0.45 对 0.42 时,这张图就在类别边界上,直接归为第一名会误导业务方,不如返回“不确定”。

from PIL import Image import torch class_names = ["cat", "dog", "rooster"] def predict(img_path): img = Image.open(img_path).convert("RGB") x = transform_val(img).unsqueeze(0) with torch.no_grad(): probs = torch.softmax(model(x), dim=1)[0] top2 = probs.argsort(descending=True)[:2] return [(class_names[i], probs[i].item()) for i in top2]

最后用一张训练时没见过的真实照片验证,而不是测试集里已经看过的图片。手机随手拍一只公鸡,缩放到 224x224,看输出概率分布是否合理。这一步能暴露训练数据和真实分布之间的差距,比如背景干扰、光线偏色、拍摄角度,都是训练集里容易缺的维度。

我做图片分类项目吃过最大的亏,是一上来就换网络结构。后来发现先把数据流程验收完,九成翻车问题都不在模型结构上。这个猫狗公鸡三分类项目,结构固定、流程完整,很适合作为你的第一个 PyTorch 实战练习。顺着代码跑通一遍,再用真实照片验证一下,你会把“模型能训练”和“模型能用”这两件事分得清清楚楚。希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询