基于PyTorch的猫狗识别:CNN、ResNet与Swin Transformer完整项目实战
2026/9/23 14:48:49 网站建设 项目流程

简介:面向机器学习与深度学习初学者、高校课程设计与毕业设计学生,提供基于Python和PyTorch的完整猫狗识别分类项目源码,涵盖CNN、ResNet、Swin Transformer等模型实现,以及数据读取、训练、测试等脚本,可帮助快速掌握图像分类的完整流程,并可直接用于课程设计、毕业设计等场景。压缩包共16个文件,含7个Python脚本、2个Markdown说明文档、1个Word设计论文、1个模型权重文件等,整体约1.67MB,结构清晰便于按需查阅。目前已有68人学习/下载。除源码外,还提供训练至400轮的CNN模型权重和日志记录,可辅助理解模型选择、训练过程与实验效果;配套说明文档与设计论文梳理了项目背景与实现思路,适合在此基础上修改扩展,完成其他分类任务。

1. 基于Python机器学习的猫狗识别分类:一套能直接跑的PyTorch完整项目

做课程设计和毕业设计的同学,多半被“猫狗识别”这四个字坑过——网上教程一大堆,但要么只给模型不给数据,要么代码缺头缺尾,跑起来全是报错。这份基于Python机器学习的猫狗识别分类项目源码,好就好在它不是一个空壳demo,而是把 CNN、ResNet、Swin Transformer 三条技术路线都塞进了同一个工程里,配套了说明文档、训练日志、模型权重和论文文档,连PyTorch环境的 events 文件都保留着。也就是说,你拿到手不是去看别人“怎么讲”,而是直接看别人“怎么跑通的”。适合三类人:急着交课程设计的学生、想复现深度学习 baseline 的初学者、以及需要一份完整代码做二次开发的从业者。它能解决的核心问题很简单:数据怎么组织、三个模型怎么训、训完怎么加载权重去做推理。

2. 工程结构和数据组织:先搞清楚 data.txt、get_data.py 和数据集目录之间怎么配合

2.1 工程里到底有哪些文件,各自干什么

拿到压缩包先别急着跑,把文件清单捋一遍比什么都重要。这个项目的文件组织很典型,是课程设计里最常见的那种“模型文件 + 工具脚本 + 说明文档”三段式结构。我拆开之后,核心文件的作用如下表:

文件/目录作用备注
get_data.py生成数据集路径清单输出 data.txt
data.txt图片路径 + 标签的文本清单训练和测试都读它
cnn.py手写 CNN 模型定义适合入门理解卷积过程
resnet.pyResNet 模型定义可能是标准 ResNet 或简化版
swin_transformer.pySwin Transformer 模型定义视觉 Transformer 路线
test_cnn.pyCNN 模型的推理脚本加载权重做单图预测
test_resnet.pyResNet 模型的推理脚本同上,针对 ResNet
show.py可视化脚本看数据、看预测结果
model/训练好的权重文件里面有 cnn_epoch400.pth
Logger/训练日志含 TensorBoard 的 events 文件
说明文档.md / Swin-trans.md使用说明和技术笔记先读这两个

这个结构最大的好处是“模型定义”和“训练推理”是解耦的。你想换模型,只需要改入口脚本,不需要动数据准备逻辑。很多初学者拿到项目就喜欢直接双击 test_cnn.py,结果报了 No module named 'cnn' 之类的错,原因就是没搞明白 Python 模块导入路径——脚本和模型文件在同一个目录,但当前工作目录不对。我一般会先在项目根目录打开终端,再执行 python test_cnn.py,而不是在 IDE 里直接右键运行。

2.2 data.txt 的格式:为什么有的项目用 ImageFolder,这个项目用文本清单

先看数据准备这一步。项目里 get_data.py 的核心职责是扫描图片目录,把每张图片的路径和类别标签写进 data.txt。常见的做法有两种:一种是用 torchvision.datasets.ImageFolder,要求数据按类别分文件夹存放;另一种就是本项目这种方式,自己维护一个“路径 标签”的文本文件。后者更灵活,因为你可以把训练集、验证集按任意比例混合,甚至可以把多个来源的图片都塞进同一个清单里。

data.txt 的每一行长这样:

data/train/cat.0.jpg 0 data/train/dog.1.jpg 1

这个01就是类别标签,0 代表猫,1 代表狗。用文本清单的好处是,你不需要复制粘贴图片到不同文件夹,只需要在生成清单时做一次路径映射。训练脚本读取这个文件后,会用 PIL 打开图片,做 resize、归一化等预处理,然后喂给模型。如果你自己想做数据增强,比如随机裁剪、水平翻转,也是在读取图片之后、送入模型之前加 transform 操作。

2.3 get_data.py 的关键逻辑:路径拼接和标签映射怎么避免踩坑

这个脚本的核心代码逻辑一般长这样:

import os # 假设你的图片按 train/cat、train/dog 存放 data_dir = "data/train" classes = ["cat", "dog"] # 类别顺序决定了标签编号 with open("data.txt", "w", encoding="utf-8") as f: for label, class_name in enumerate(classes): class_dir = os.path.join(data_dir, class_name) for img_name in os.listdir(class_dir): if img_name.endswith((".jpg", ".jpeg", ".png")): img_path = os.path.join(class_dir, img_name) f.write(f"{img_path} {label}\n")

这个脚本的逻辑很简单:遍历每个类别文件夹,把所有图片的路径和对应的数字标签写进 data.txt。注意classes列表的顺序很重要——如果你把["dog", "cat"]写在前面,那 dog 就变成 0,猫就变成 1,后面训练出来的模型语义就完全反了。我在实际跑项目时习惯打印 data.txt 的前几行确认标签正确,这一步花不了十秒钟,但能省掉后面排查预测结果“猫狗颠倒”的半天时间。

另外,路径分隔符要注意。Windows 下 os.path.join 生成的是反斜杠\,而 Linux 下是正斜杠/。如果你的训练脚本跑在 Linux 服务器上,data.txt 里却是 Windows 路径,会直接 FileNotFoundError。我一般会在 get_data.py 里把所有路径统一替换成正斜杠:

img_path = os.path.join(class_dir, img_name).replace("\\", "/")

这个小改动,能让你的项目在 Windows 写完、Linux 上跑的时候不翻车。

3. 三条模型路线对比:CNN、ResNet、Swin Transformer 各自怎么选

3.1 手写 CNN:最适合理解卷积本质,也最容易过拟合

cnn.py 里定义的是一个手工搭建的卷积神经网络。课程设计里最常见的手写 CNN 结构是“卷积 + 池化 + 全连接”的堆叠。一个典型的代码框架如下:

import torch.nn as nn class SimpleCNN(nn.Module): def __init__(self, num_classes=2): super(SimpleCNN, self).__init__() # 第一个卷积块:3 通道输入,16 个卷积核,3x3 大小 self.conv1 = nn.Sequential( nn.Conv2d(3, 16, kernel_size=3, padding=1), nn.ReLU(), nn.MaxPool2d(2) # 尺寸减半 ) # 第二个卷积块:16 -> 32 通道 self.conv2 = nn.Sequential( nn.Conv2d(16, 32, kernel_size=3, padding=1), nn.ReLU(), nn.MaxPool2d(2) ) # 全连接分类头:输入维度取决于最后的特征图尺寸 self.fc = nn.Linear(32 * 56 * 56, num_classes) def forward(self, x): x = self.conv1(x) x = self.conv2(x) # 展平后送入全连接层 x = x.view(x.size(0), -1) x = self.fc(x) return x

这里的nn.Conv2d(3, 16, kernel_size=3, padding=1)表示输入是 3 通道(RGB),输出 16 个特征图,卷积核 3x3,padding 为 1 保持尺寸不变。nn.MaxPool2d(2)会把特征图宽高各缩一半。如果输入图片是 224x224,经过两次池化变成 56x56,所以全连接层的输入维度是32 * 56 * 56。这个数字是硬算出来的,不是随便写的——如果你把输入图片尺寸改成 128x128,这里必须同步调整,否则会报维度不匹配的错误。

手写 CNN 的优势是可控性强,每一层的输出你都可以打印出来看,非常适合写论文时画网络结构图。但缺点也明显:在猫狗识别这种任务上,自己搭的 CNN 很容易过拟合——训练集准确率 98%,验证集只有 80%。原因在于参数量虽然不大,但特征提取能力不够强,学不到足够泛化的语义特征。如果你发现训练 loss 降得很低但验证 loss 很高,优先考虑加 Dropout、做数据增强、或者直接换 ResNet。

3.2 ResNet:残差结构解决退化问题,是精度和速度的平衡点

resnet.py 里实现的是带残差连接的 ResNet。残差连接的核心思想是:与其让网络直接学习一个复杂的映射 H(x),不如让它学习残差 F(x) = H(x) - x,然后通过跳跃连接把输入 x 加到输出上。这样做的好处是梯度可以绕过中间的卷积层直接回传,解决了网络加深后的退化问题。

一个简化的残差块实现如下:

import torch.nn as nn class ResidualBlock(nn.Module): def __init__(self, in_channels, out_channels, stride=1): super(ResidualBlock, self).__init__() self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1, bias=False) self.bn1 = nn.BatchNorm2d(out_channels) self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False) self.bn2 = nn.BatchNorm2d(out_channels) # 如果输入输出通道数不一致,用 1x1 卷积对齐 self.shortcut = nn.Sequential() if stride != 1 or in_channels != out_channels: self.shortcut = nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride, bias=False), nn.BatchNorm2d(out_channels) ) def forward(self, x): identity = self.shortcut(x) out = torch.relu(self.bn1(self.conv1(x))) out = self.bn2(self.conv2(out)) out += identity # 残差连接 return torch.relu(out)

注意这里的shortcut分支。当 stride=2 或者通道数变化时,输入输出形状不一致,必须用 1x1 卷积把 x 的通道数和尺寸对齐,否则out += identity会直接报错。这个细节是 ResNet 实现里最容易错的地方,很多人从 GitHub 抄代码跑不通,问题往往出在这里。

在猫狗识别任务上,ResNet 的表现明显好于手写 CNN。原因是残差结构让网络可以训练得更深,更深意味着能提取到更高层的语义特征——比如狗耳朵的形状、猫胡须的纹理。项目里 resnet.py 对应的测试脚本是 test_resnet.py,说明作者最终的主推模型大概率是 ResNet 路线。

3.3 Swin Transformer:视觉 Transformer 路线,适合想冲高精度的场景

swin_transformer.py 是三个模型里最复杂的一个。Swin Transformer 的核心是移动窗口自注意力(Shifted Window Attention),它把图片划分成固定大小的窗口,在窗口内做自注意力计算,然后在相邻层之间偏移窗口,让不同窗口之间的信息可以交互。这样既保留了 Transformer 的全局建模能力,又把计算复杂度从 O(N²) 降到了窗口级别的可控范围。

Swin Transformer 的代码量比 CNN 大一个量级,里面涉及窗口划分、相对位置编码、Patch Merging 等一堆操作。对课程设计来说,你不需要从零默写整个模型——直接调用项目里现成的 swin_transformer.py 就行,但你要能说清楚它和 CNN 的本质区别:CNN 是局部感受野的卷积操作堆叠,Swin 是通过自注意力机制建模像素之间的长距离依赖。

如果你要在论文里对比三个模型的精度,结论一般是:Swin Transformer > ResNet > 手写 CNN,但训练时间和显存占用也是倒过来的。Swin 对 GPU 显存的要求最高,如果显卡只有 4GB 显存,batch size 大概率只能设到 8 甚至 4,训练速度会让你怀疑人生。我个人的建议是:如果毕设要求“创新点”,你可以把 Swin Transformer 作为主模型,用小数据集微调,再拿 ResNet 做 baseline 对比,这样既有工作量又有深度。

3.4 三份训练日志怎么读:epoch 400 的 CNN 告诉你什么信息

项目里保留了一个cnn_epoch400.pth的权重文件,这暗示作者用 CNN 训练了整整 400 轮。看到这个数字,有经验的人第一反应是:要么作者用了早停(early stopping)但初始设置轮数很充裕,要么模型在小数据集上反复震荡。深度学习里有一个不成文的经验:训练轮数翻倍,不代表精度的提升能翻倍,很多时候从第 200 轮到第 400 轮,验证集准确率可能只涨了 1% 到 2%。

训练日志目录里还有一个 TensorFlow 的 events 文件——events.out.tfevents.1649819535...。这就比较有趣了,说明作者在开发过程中既用过 PyTorch 也摸过 TensorFlow,或者 TensorBoard 的日志格式残留。如果你想在本地复现训练过程的可视化,可以用 TensorBoard 加载这个目录看 loss 曲线;不想折腾的话,直接打开 Logger 里的文本日志看 loss 数值变化也够用。

4. 训练和推理实战:从加载权重到单图预测,把整个流程跑通

4.1 test_cnn.py 的推理逻辑:权重文件怎么加载,输入怎么预处理

现在进入最关键的环节——用训练好的权重做预测。先看 test_cnn.py 的典型流程:

import torch import torchvision.transforms as transforms from PIL import Image from cnn import SimpleCNN # 1. 定义与训练时完全一致的预处理 transform = transforms.Compose([ transforms.Resize((224, 224)), # 尺寸必须和训练时一致 transforms.ToTensor(), # HWC -> CHW,像素值归一化到 [0,1] transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) # ImageNet 统计量 ]) # 2. 加载模型和权重 model = SimpleCNN(num_classes=2) checkpoint = torch.load("model/cnn_epoch400.pth", map_location="cpu") model.load_state_dict(checkpoint["model_state_dict"] if "model_state_dict" in checkpoint else checkpoint) model.eval() # 3. 推理单张图片 def predict(image_path): img = Image.open(image_path).convert("RGB") # 强制转 RGB img_tensor = transform(img).unsqueeze(0) # 增加 batch 维度 with torch.no_grad(): outputs = model(img_tensor) _, predicted = torch.max(outputs, 1) return "猫" if predicted.item() == 0 else "狗" print(predict("test_cat.jpg"))

这里有几个细节必须注意。第一,map_location="cpu"是给没有 GPU 的机器用的,如果你有 CUDA 且权重是在 GPU 上训练的,可以去掉这个参数或者改成map_location="cuda:0"。第二,model.eval()必须调用——它会关闭 Dropout 和 BatchNorm 的训练行为,否则同样的输入每次预测结果可能不一样,而且精度会下降。第三,Image.open(...).convert("RGB")很重要,因为有些图片是 RGBA 四通道或者灰度单通道,不转成 RGB 的话,输入通道数和模型定义不匹配,会直接报错。

4.2 权重文件是完整 checkpoint 还是纯 state_dict?load 的时候怎么写

这是推理脚本里最容易出幺蛾子的地方。torch.save 有两种常见姿势:一种是只保存模型参数(state_dict),另一种是保存包含优化器状态、epoch 信息在内的完整 checkpoint。项目里的权重文件是cnn_epoch400.pth,从命名看很可能是完整 checkpoint。所以我在代码里做了兼容处理:

if "model_state_dict" in checkpoint: model.load_state_dict(checkpoint["model_state_dict"]) else: model.load_state_dict(checkpoint)

这段代码的意思是:先检查字典里有没有model_state_dict这个键,如果有就用它;如果没有,说明直接保存的就是 state_dict。这种写法能兼容两种情况,不管作者当初是用torch.save(model.state_dict(), ...)还是torch.save({"model_state_dict": model.state_dict(), ...}, ...)保存的,都能正确加载。我拆过好几个开源项目,发现很多作者保存权重的习惯都不一样,所以这个“双保险”写法值得养成习惯。

4.3 训练入口复现:train 脚本缺失时,怎么补一个最简训练流程

压缩包里我注意到没有明确列出 train.py,但既然有模型定义和权重文件,训练流程大概率是作者在 notebook 或临时脚本里跑的。如果你想自己重训模型,可以用下面这个最简训练循环:

import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import Dataset, DataLoader from PIL import Image from cnn import SimpleCNN # 自定义 Dataset,读取 data.txt class CatDogDataset(Dataset): def __init__(self, txt_path, transform=None): self.samples = [] self.transform = transform with open(txt_path, "r", encoding="utf-8") as f: for line in f: path, label = line.strip().split() self.samples.append((path, int(label))) def __len__(self): return len(self.samples) def __getitem__(self, idx): path, label = self.samples[idx] img = Image.open(path).convert("RGB") if self.transform: img = self.transform(img) return img, label # 训练参数 batch_size = 32 learning_rate = 0.001 epochs = 50 # 初始化 dataset = CatDogDataset("data.txt", transform=transform) dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=True, num_workers=0) model = SimpleCNN() criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=learning_rate) # 训练循环 for epoch in range(epochs): running_loss = 0.0 for images, labels in dataloader: optimizer.zero_grad() outputs = model(images) loss = criterion(outputs, labels) loss.backward() optimizer.step() running_loss += loss.item() print(f"Epoch {epoch+1}/{epochs}, Loss: {running_loss/len(dataloader):.4f}")

nn.CrossEntropyLoss()是分类任务默认的损失函数,它内部已经包含了 softmax 操作,所以模型最后一层不需要额外加 softmax。optimizer.zero_grad()必须在每次反向传播前清空梯度,否则梯度会累加导致训练不稳定。num_workers=0是 Windows 环境下的安全设置,写成大于 0 的值在 Windows 上经常报多进程相关的错。

4.4 show.py 是干什么的:可视化预测结果,论文插图就靠它

show.py 这个脚本通常负责把预测结果可视化——把图片读进来,跑一次模型,然后在图片上画一个标题框,写着“猫”或“狗”,最后保存成一张带标注的图。这个脚本对论文和答辩 PPT 特别有用。你可以用它批量处理几张典型图片,生成“模型正确识别”“模型错误分类”的对比图,放进论文的实验分析章节,比纯文字描述有说服力得多。

5. 避坑指南:从路径到显存,五个最常见的翻车现场

5.1 现象:data.txt 里的路径在 Windows 上能跑,换到 Linux 就 FileNotFoundError

原因很简单:Windows 路径分隔符是反斜杠,Linux 是正斜杠。get_data.py 生成 data.txt 时用的是 os.path.join,在 Windows 上自然生成反斜杠路径。你把项目传到服务器上用 Linux 跑,Python 会老老实实地把反斜杠当成文件名的一部分,自然找不到文件。

解决办法是在 get_data.py 里强制替换所有分隔符:

img_path = os.path.join(class_dir, img_name).replace("\\", "/")

或者更稳妥地在训练脚本读取 data.txt 时做一次统一化处理:

from pathlib import Path path, label = line.strip().split() path = str(Path(path)) # 自动转换为当前系统的标准路径格式

从那以后我每次生成数据集清单,都会先跑一遍脚本然后head -n 5 data.txt看一眼路径格式,这个习惯帮我避开了不少跨平台翻车的坑。

5.2 现象:加载权重时报错 size mismatch for fc.weight: copying a param with shape torch.Size([2, 512]) from checkpoint

原因是权重文件对应的模型结构和当前代码定义的模型结构不一致。常见情况有三种:一是你改了模型定义里的num_classes,比如从 2 改成了 10;二是全连接层之前的特征维度不对,比如输入图片尺寸变了导致展平后的维度变了;三是作者训练用的 ResNet 和你代码里定义的 ResNet 层数不一样。

解决办法是对比权重文件的键名和当前模型的键名,找出差异在哪一层。可以在加载前打印出来:

checkpoint = torch.load("model/cnn_epoch400.pth", map_location="cpu") model = SimpleCNN() print("Checkpoint keys:", list(checkpoint.keys())[:5]) print("Model keys:", list(model.state_dict().keys())[:5])

把两边的形状逐一对比,就能定位是哪个层的维度对不上。如果是num_classes不一致,那就别纠结,直接改模型定义或者在加载时把最后一层剥离。

5.3 现象:inference 时同样的图片,每次预测结果不一样,准确率还不稳定

这个现象十有八九是忘了调用model.eval()。训练模式下 BatchNorm 层会使用当前 batch 的均值和方差,Dropout 层会随机失活一部分神经元,所以模型在训练模式下的前向传播是有随机性的。推理时必须切换到 eval 模式,BatchNorm 才会使用训练时积累的全局统计量,Dropout 才会全部保留。

另外还要注意,如果你对输入图片做了随机预处理,比如随机裁剪、随机翻转,那推理结果当然每次都不一样,所以要保证预处理函数里没有随机操作。

5.4 现象:训练时 loss 一直不降,或者直接变成 NaN

先说 loss 是 NaN 的情况:最常见的原因是学习率设置过大,导致梯度爆炸。CNN 模型建议从 0.001(Adam 优化器)开始试,如果 loss 在几个 epoch 内就冲到 NaN,果断把学习率除以 10。还有一种可能是输入数据里有损坏的图片,PIL 打开失败返回空数组,喂给模型后出现异常数值。排查方法是在 Dataset 的__getitem__里加 try-except,把加载失败的图片路径打印出来。

loss 一直不降的情况则往往和数据预处理有关。比如图片没有做归一化,像素值范围是 0 到 255 而不是 0 到 1,这时候梯度传播的尺度就很奇怪,模型很难收敛。确认你的 transform 里加了transforms.ToTensor(),它会自动把像素值从 0-255 缩放到 0.0-1.0。

5.5 现象:CUDA out of memory 在 Swin Transformer 上频繁出现

Swin Transformer 的显存占用是三个模型里最高的,视野注意力机制的中间张量很大,BatchNorm 层的缓存也占空间。如果你显卡只有 4GB 显存,跑 Swin 会很吃力。解决办法按优先级排列:先减小 batch_size,比如从 16 减到 8 或 4;再考虑减小输入图片尺寸,比如从 224 改成 192;最后可以把混合精度训练打开——PyTorch 自带torch.cuda.amp可以自动用半精度浮点数计算,显存占用直接砍半。混精训练的代码改动不大,核心就三行:

scaler = torch.cuda.amp.GradScaler() with torch.cuda.amp.autocast(): outputs = model(images) loss = criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()

如果做完这些还是爆显存,那只能说明硬件不合适跑 Swin,老老实实回去用 ResNet,别硬扛。

6. 验证模型是否靠谱:用混淆矩阵和真实样例做最后的把关

6.1 不要只看 Accuracy,按类别拆开看混淆矩阵

课程设计答辩的时候,老师说“你的模型准确率 95% 挺高的”,但一调出测试结果发现狗全部识别正确、猫有一半被认成狗——这种情况我见过太多次了。二分类任务的准确率非常具有迷惑性,如果测试集里狗占了 80%、猫占 20%,模型只要无脑全猜狗就能拿到 80% 准确率。所以验证时一定要看每个类别分别的召回率,也就是混淆矩阵。

用 PyTorch 统计混淆矩阵,核心代码很直观:

import torch def compute_confusion_matrix(model, dataloader, device="cpu"): model.eval() model.to(device) # 2x2 矩阵: [真实猫][真实狗] x [预测猫][预测狗] cm = torch.zeros(2, 2, dtype=torch.int64) with torch.no_grad(): for images, labels in dataloader: images, labels = images.to(device), labels.to(device) outputs = model(images) _, preds = torch.max(outputs, 1) for t, p in zip(labels.view(-1), preds.view(-1)): cm[t, p] += 1 return cm cm = compute_confusion_matrix(model, test_loader) print("混淆矩阵:") print(cm) # 每一行代表真实类别,每一列代表预测类别

这个脚本会输出一个 2x2 的矩阵,对角线上的数字是正确分类的数量,反对角线上的数字是错误分类的数量。如果cm[0][1]明显大于cm[1][0],说明模型把很多猫错认成了狗——这时候常见解决办法是给训练数据里的猫做更多的数据增强,或者收集更多猫的图片来平衡两个类别。

6.2 用真实场景图片做边界测试,模型有没有泛化能力一看便知

除了用测试集算指标,我强烈建议你从网上找几张不在这份数据集里的猫狗图片,特别是那种“角度刁钻”的——比如只露半个猫头、狗在奔跑中、白色猫在白色背景下。把这些图片喂给 test_cnn.py 跑一遍,你会很快发现模型的真实水平。因为训练集里的图片通常是正对镜头、主体居中的标准图,而真实场景的图片各种姿势都有,泛化能力弱的模型在这类图上会原形毕露。

如果你想更系统地测,可以准备一个“挑战集”,里面放十张不同类型图片:卡通猫、布偶猫、柯基、金毛、黑猫夜间照等。一张一张预测,把正确和错误的结果记录下来。这个过程会直接暴露模型学到的到底是“猫狗的特征”还是“数据集的背景特征”。比如模型把所有带白色背景的图都预测成狗,那它大概率就是学到了数据集的偏差,而不是真正的语义特征。

6.3 训练曲线怎么看:过拟合的早期信号是 gap 拉大而不是 loss 升高

最后一个技巧是训练曲线解读。很多人只看最终准确率,不看训练过程。正确的做法是同时盯着训练集 loss 和验证集 loss 两条曲线。当训练集 loss 还在下降、但验证集 loss 开始回升或不再下降时,这个“gap 拉大”的时刻就是过拟合的起点。项目里的 Logger 目录既然保留了 TensorBoard 的 events 文件,我建议你本地起一下 TensorBoard 看看原始曲线:

tensorboard --logdir Logger/

如果你的环境没装 TensorBoard,直接看文本日志里的 loss 数值变化也行——发现训练 loss 降到 0.01 以下而验证 loss 还在 0.5 以上,说明模型已经死记硬背了训练集,这时候最有效的干预不是继续训,而是降低模型复杂度或者加强数据增强。从那以后我每次拿到别人的模型权重,第一件事不是跑准确率,而是先看它的训练曲线和混淆矩阵,确认这份权重是真的训出来的、不是运气好碰出来的,再决定要不要在自己的场景里用。希望这份项目拆解能帮你把猫狗识别这条技术路线真正跑通,少走那些我已经替你踩过的弯路。

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

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

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

立即咨询