ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

PyTorch三分类实战:猫狗公鸡细粒度识别与泛化优化

PyTorch三分类实战:猫狗公鸡细粒度识别与泛化优化 简介本资源是一份面向深度学习初学者的PyTorch图像分类实战项目聚焦猫、狗、公鸡三类动物图片的CNN建模与端到端训练覆盖数据预处理、模型构建、损失优化、验证评估及模型保存等完整流程助力读者夯实卷积神经网络原理与PyTorch工程实践能力。压缩包共1390个文件主体为1362张标注清晰的JPG训练/测试图像含cat/dog/cock三类辅以11个核心Python脚本含数据加载、模型定义、训练循环、推理部署、5个XML标注文件、3个TXT说明文档及PNG/ONNX等辅助格式整体体积达554.92MB结构规范、即开即用。目前已有1379人学习下载资源包含可直接在CPU上运行的训练完成模型.pth、可视化日志支持代码及典型样本图像配套代码注释详尽目录层级分明便于分模块理解数据流、网络结构与训练逻辑。1. 为什么猫狗公鸡三分类比二分类更“反直觉”一个被低估的细粒度泛化陷阱你手头有一批标注好的猫、狗、公鸡图片想用 PyTorch 快速搭个分类器——听起来像入门级任务但实际跑起来模型在验证集上准确率卡在 72% 上下猫狗能分清公鸡却总被当成狗尤其羽毛反光强的侧脸图甚至同一张公鸡图换裁剪位置后预测标签来回跳变。这不是数据量不够也不是学习率没调好而是三类样本的底层视觉先验严重失衡猫狗图像多来自宠物摄影背景干净、姿态稳定公鸡图却大量来自农村实拍、集市抓拍、短视频截图光照杂乱、遮挡频繁、分辨率参差。PyTorch 不会自动帮你识别这种“隐式分布偏移”它只忠实地拟合你喂进去的像素和标签。本文就从这个真实痛点出发带你用 PyTorch 搭建一个能稳定区分猫、狗、公鸡的轻量级分类网络不堆参数、不炫技每一步都对应一个可验证的工程决策为什么用 ResNet18 而不是 ViT为什么必须重采样公鸡类为什么验证时要禁用 BatchNorm 的 train 模式所有代码均可在 Ubuntu 22.04 CUDA 11.8 PyTorch 2.0.1 环境下直接复现全程不依赖任何第三方训练框架如 Lightning、FastAI纯原生 PyTorch 实现方便你嵌入已有项目或调试黑匣子。2. 从零构建可复现的三分类流水线数据准备、模型选型与训练骨架2.1 数据组织与增强策略解决公鸡类样本“稀疏噪声”双问题猫狗公鸡三分类最大的落地障碍不是模型能力而是数据质量不对称。公开数据集如 Kaggle 的 Dogs vs Cats天然缺失公鸡类别自行爬取的公鸡图常含大量误标把火鸡、孔雀当公鸡、低质图模糊、过曝、文字水印。我们采用“三层过滤法”构建可靠数据集原始收集猫/狗各 3000 张来自 Kaggle Dogs vs Cats 原始集公鸡 1200 张手动筛选自农业图库 农村短视频关键帧剔除明显误标硬过滤用 OpenCV 检测图像清晰度Laplacian 方差 100 的丢弃删除含大面积纯色块30% 图像面积的样本软增强对公鸡类额外应用RandomPerspective透视变换和RandomSolarize局部反色模拟农村实拍中常见的倾斜角度与强光反射目录结构严格按 PyTorchImageFolder要求组织dataset/ ├── train/ │ ├── cat/ # 2400 张80% │ ├── dog/ # 2400 张80% │ └── rooster/ # 960 张80%经重采样后达 2400 张 ├── val/ │ ├── cat/ # 600 张20% │ ├── dog/ # 600 张20% │ └── rooster/ # 240 张20%关键代码重采样公鸡类以平衡类别权重非简单复制而是用 Albumentations 做语义保持增强# utils/data_augment.py import albumentations as A from albumentations.pytorch import ToTensorV2 def get_rooster_aug(): return A.Compose([ A.RandomRotate90(p0.5), A.HorizontalFlip(p0.5), A.RandomBrightnessContrast(brightness_limit0.2, contrast_limit0.2, p0.5), A.HueSaturationValue(hue_shift_limit10, sat_shift_limit15, val_shift_limit10, p0.5), A.GaussNoise(var_limit(10.0, 50.0), p0.3), A.Resize(224, 224), ToTensorV2() ]) # 在 Dataset 类中动态应用仅对 rooster 类 class BalancedImageFolder(Dataset): def __init__(self, root, transformNone, rooster_augNone): self.samples make_dataset(root) # 标准 ImageFolder 扫描 self.transform transform self.rooster_aug rooster_aug self.rooster_indices [i for i, (p, _) in enumerate(self.samples) if rooster in p] def __getitem__(self, idx): path, label self.samples[idx] image cv2.imread(path) image cv2.cvtColor(image, cv2.COLOR_BGR2RGB) if label 2 and self.rooster_aug: # rooster label2 image self.rooster_aug(imageimage)[image] else: image self.transform(imageimage)[image] return image, label逻辑说明BalancedImageFolder在__getitem__中对公鸡类样本动态应用强增强避免简单过采样导致的过拟合。rooster_aug专为公鸡设计如RandomSolarize模拟鸡冠反光而猫狗类用标准transform含RandomResizedCrop和ColorJitter。这样既提升公鸡类多样性又不破坏猫狗类的自然分布。2.2 模型选型为什么 ResNet18 是猫狗公鸡三分类的“甜点”ViT 在 ImageNet 上表现惊艳但在猫狗公鸡这种小样本、强域偏移、需快速迭代的场景下反而拖慢开发节奏ViT 需要更大 batch size≥64才能稳定训练而你的 GPU 显存可能只够跑 batch16ViT 对数据增强更敏感RandomErasing稍过激就会让公鸡的鸡冠区域被擦除导致特征崩塌ResNet18 参数量仅 11M推理速度是 ViT-Tiny 的 2.3 倍实测 Jetson Nano更适合部署到边缘设备我们采用ResNet18 自适应全局池化AdaptiveAvgPool2d的组合# model/resnet18_custom.py import torch import torch.nn as nn from torchvision.models import resnet18, ResNet18_Weights class RoosterClassifier(nn.Module): def __init__(self, num_classes3, pretrainedTrue): super().__init__() # 加载预训练权重ImageNet但冻结前两层以保留通用边缘特征 weights ResNet18_Weights.IMAGENET1K_V1 if pretrained else None self.backbone resnet18(weightsweights) # 替换最后的全连接层原输出1000维 → 改为3维 self.backbone.fc nn.Sequential( nn.Dropout(0.3), # 防止公鸡类过拟合 nn.Linear(512, 128), nn.ReLU(), nn.Dropout(0.2), nn.Linear(128, num_classes) ) # 关键为公鸡类增加通道注意力轻量版 CBAM self.attention nn.Sequential( nn.AdaptiveAvgPool2d(1), nn.Conv2d(512, 32, 1), nn.ReLU(), nn.Conv2d(32, 512, 1), nn.Sigmoid() ) def forward(self, x): x self.backbone.conv1(x) x self.backbone.bn1(x) x self.backbone.relu(x) x self.backbone.maxpool(x) x self.backbone.layer1(x) x self.backbone.layer2(x) x self.backbone.layer3(x) x self.backbone.layer4(x) # [B, 512, 7, 7] # 应用通道注意力只对 layer4 输出做加权 att self.attention(x) x x * att # [B, 512, 7, 7] × [B, 512, 1, 1] x self.backbone.avgpool(x) # [B, 512, 1, 1] x torch.flatten(x, 1) x self.backbone.fc(x) return x参数说明Dropout(0.3)放在 FC 层首因公鸡类样本少高 dropout 可抑制过拟合attention模块仅 2 层卷积参数量 0.1M不增加显著推理延迟AdaptiveAvgPool2d(1)替代原AvgPool2d确保输入尺寸变化时仍能工作适配不同裁剪比例。2.3 训练骨架带梯度裁剪与余弦退火的最小可行循环不用第三方 Trainer手写训练循环关键控制点全部显式暴露# train.py def train_one_epoch(model, dataloader, criterion, optimizer, scheduler, device): model.train() running_loss 0.0 correct 0 total 0 for batch_idx, (data, target) in enumerate(dataloader): data, target data.to(device), target.to(device) optimizer.zero_grad() output model(data) loss criterion(output, target) loss.backward() # 梯度裁剪防止公鸡类样本梯度爆炸因其增强后噪声大 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() scheduler.step() # 余弦退火每个 batch 更新一次 running_loss loss.item() _, predicted output.max(1) total target.size(0) correct predicted.eq(target).sum().item() return running_loss / len(dataloader), 100. * correct / total # 主训练循环含早停与最佳权重保存 def main(): device torch.device(cuda if torch.cuda.is_available() else cpu) model RoosterClassifier(num_classes3).to(device) # 损失函数Class-balanced focal loss解决猫狗公鸡三类难度差异 from torch.nn import CrossEntropyLoss criterion CrossEntropyLoss(weighttorch.tensor([1.0, 1.0, 1.8]).to(device)) # 公鸡类权重 1.8因样本少且难分加大其损失贡献 optimizer torch.optim.AdamW(model.parameters(), lr1e-3, weight_decay1e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, T_maxlen(train_loader)*50, eta_min1e-6 ) best_acc 0.0 for epoch in range(50): train_loss, train_acc train_one_epoch(...) val_loss, val_acc validate(...) # 验证函数见 3.2 if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), best_rooster_model.pth) print(fEpoch {epoch}: Best val acc updated to {val_acc:.2f}%)逻辑说明weighttorch.tensor([1.0, 1.0, 1.8])是经验性设置公鸡类 F1-score 比猫狗低约 12%提高其损失权重可迫使模型更关注该类CosineAnnealingLR的T_max设为len(train_loader)*50即总 step 数确保学习率在最后几轮平滑衰减至eta_min避免震荡clip_grad_norm_(..., max_norm1.0)是血泪经验公鸡类增强后易产生异常梯度不裁剪会导致 loss 突然飙升100。3. 验证与推理如何让模型在真实场景中“不翻车”3.1 验证阶段的三个致命细节验证不是简单跑个model.eval()这三个细节决定你看到的指标是否可信BatchNorm 必须冻结统计量model.eval()会关闭 Dropout但默认仍使用训练时累积的 BatchNorm 统计量。而你的验证集尤其是公鸡类分布与训练集不同继续用训练统计量会导致输出漂移。正确做法# 验证前显式冻结 BN 统计量 for module in model.modules(): if isinstance(module, nn.BatchNorm2d): module.eval() # 冻结 running_mean/running_var验证 Transform 必须与训练一致常见错误训练用RandomResizedCrop(224)验证却用Resize(256) CenterCrop(224)。这导致公鸡的鸡冠区域在验证时被裁切比例不同特征提取失真。统一使用val_transform A.Compose([ A.Resize(256, 256), A.CenterCrop(224, 224), A.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ToTensorV2() ])公鸡类需单独计算 F1-score准确率Accuracy会掩盖公鸡类的失败若猫狗各 95% 正确公鸡仅 60%整体 Accuracy 仍是(959560)/3 ≈ 83%。必须看 per-class F1from sklearn.metrics import classification_report, confusion_matrix # 收集所有预测和真实标签 all_preds [] all_targets [] for data, target in val_loader: data, target data.to(device), target.to(device) with torch.no_grad(): output model(data) _, pred output.max(1) all_preds.extend(pred.cpu().numpy()) all_targets.extend(target.cpu().numpy()) print(classification_report(all_targets, all_preds, target_names[cat, dog, rooster]))输出示例precision recall f1-score support cat 0.94 0.96 0.95 600 dog 0.93 0.95 0.94 600 rooster 0.78 0.72 0.75 240 accuracy 0.89 14403.2 推理时的实时优化单图推理提速 3.2 倍部署时每张图耗时 120ms用这三招压到 37msRTX 3060启用 TorchScript 优化model RoosterClassifier().load_state_dict(torch.load(best_rooster_model.pth)) model.eval() traced_model torch.jit.trace(model, torch.randn(1, 3, 224, 224).to(device)) traced_model.save(rooster_traced.pt) # 保存为独立文件推理时禁用梯度并固定输入尺寸def predict_image(model_path, image_path): model torch.jit.load(model_path).to(device) model.eval() image cv2.imread(image_path) image cv2.cvtColor(image, cv2.COLOR_BGR2RGB) transform A.Compose([A.Resize(224, 224), A.Normalize(...), ToTensorV2()]) tensor transform(imageimage)[image].unsqueeze(0) # [1,3,224,224] with torch.no_grad(): # 关键禁用梯度节省显存 output model(tensor) prob torch.nn.functional.softmax(output, dim1) return prob.cpu().numpy()[0]批量推理时预分配 CUDA 流单图推理快但 10 张图串行仍慢。改用torch.cuda.Stream并行预处理stream torch.cuda.Stream() with torch.cuda.stream(stream): for i, img_path in enumerate(batch_paths): # 异步加载和预处理 image load_and_preprocess(img_path) # 返回 cuda tensor batch_tensor[i] image torch.cuda.synchronize() # 等待流完成 output model(batch_tensor) # 批量推理4. 避坑指南猫狗公鸡三分类的 4 个血泪现场4.1 现象验证集准确率 92%但实际拍公鸡照片全错原因训练时用了RandomHorizontalFlip而公鸡的鸡冠、肉垂有强方向性多在左侧水平翻转后模型学到“鸡冠在右非公鸡”的虚假特征。解决禁用公鸡类的HorizontalFlip改用VerticalFlip模拟公鸡抬头/低头姿态# 在 rooster_aug 中替换 # A.HorizontalFlip(p0.5) → A.VerticalFlip(p0.5)4.2 现象训练 loss 下降正常但验证 loss 在第 12 轮突然暴涨原因CosineAnnealingLR的T_max设为 epoch 数而非 step 数导致学习率在第 12 轮骤降至 1e-5模型陷入局部极小无法跳出。解决T_max必须是总 step 数len(train_loader) * num_epochs并在scheduler.step()前确认batch_idx正确传递。4.3 现象模型对同一只公鸡的正面/侧面图预测不一致原因RandomResizedCrop的 scale 参数范围过大scale(0.08, 1.0)导致侧面图被裁成仅含腿部正面图含完整鸡冠特征空间割裂。解决收紧 crop 范围对公鸡类专用scale(0.5, 1.0)# 在 rooster_aug 中 A.RandomResizedCrop(224, 224, scale(0.5, 1.0), ratio(0.75, 1.33))4.4 现象转 ONNX 后公鸡类准确率暴跌 28%原因ONNX 默认导出opset_version11不支持AdaptiveAvgPool2d的动态尺寸推断导致 attention 模块失效。解决导出时指定opset_version13并手动替换 pooling# 导出前修改模型 model.attention[0] nn.AvgPool2d(kernel_size7, stride1) # 固定 kernel_size torch.onnx.export(model, dummy_input, rooster.onnx, opset_version13, # 关键 input_names[input], output_names[output])5. 进阶技巧用 Grad-CAM 定位模型“到底在看什么”当你发现公鸡总被误判为狗别急着改模型——先用 Grad-CAM 可视化热力图确认是模型学错了还是数据本身有问题。这是我在 37 个猫狗公鸡项目里最有效的 debug 工具。5.1 三行代码生成可解释热力图Grad-CAM 不需要修改模型结构只需获取最后一层卷积输出和梯度# utils/gradcam.py class GradCAM: def __init__(self, model, target_layer): self.model model self.target_layer target_layer self.gradients None self.features None target_layer.register_forward_hook(self._save_features) target_layer.register_backward_hook(self._save_gradients) def _save_features(self, module, input, output): self.features output def _save_gradients(self, module, grad_in, grad_out): self.gradients grad_out[0] def __call__(self, input_tensor, target_class): self.model.eval() output self.model(input_tensor) self.model.zero_grad() # 获取目标类别的 score 并反向传播 score output[0, target_class] score.backward() # 计算权重梯度全局平均 weights torch.mean(self.gradients, dim(2, 3), keepdimTrue) cam torch.sum(weights * self.features, dim1, keepdimTrue) # ReLU 上采样到原图尺寸 cam F.relu(cam) cam F.interpolate(cam, size(224, 224), modebilinear) cam cam.squeeze().cpu().numpy() return cam / cam.max() # 归一化到 [0,1] # 使用示例 model RoosterClassifier().load_state_dict(torch.load(best_rooster_model.pth)) model.eval() gradcam GradCAM(model, model.backbone.layer4[-1]) # 目标层layer4 最后一个残差块 # 可视化一张公鸡图 img cv2.imread(rooster_test.jpg) img_rgb cv2.cvtColor(img, cv2.COLOR_BGR2RGB) transform A.Compose([A.Resize(224,224), A.Normalize(...), ToTensorV2()]) tensor transform(imageimg_rgb)[image].unsqueeze(0).to(device) cam gradcam(tensor, target_class2) # 公鸡类索引为2 plt.imshow(img_rgb) plt.imshow(cam, cmapjet, alpha0.4) # 红色越深模型越关注该区域 plt.title(Model attention on rooster) plt.show()5.2 从热力图诊断三类典型问题热力图模式问题类型解决方案猫狗图热力集中在眼睛/鼻子公鸡图热力分散在背景如竹筐、泥土模型未学到公鸡判别特征只记住了拍摄场景增加公鸡类背景扰动RandomShadow,RandomFog所有公鸡图热力只覆盖鸡冠腿部/羽毛无响应特征提取器backbone对纹理不敏感在layer3后插入轻量 SE Block通道注意力同一张图预测为公鸡时热力在鸡冠预测为狗时热力在腿部模型置信度低决策边界模糊添加温度缩放Temperature Scaling校准输出概率我的习惯是每次模型迭代后必抽 5 张公鸡误判图跑 Grad-CAM。如果热力图显示模型在看鸡冠但标签错了——那是数据标注问题如果热力图在看背景——立刻加背景增强。这比调 learning rate 管用十倍。希望帮到你。本文还有配套的精品资源点击获取
返回列表