资讯动态

CNN鸟类识别全流程:PyTorch训练+Flask部署+微信小程序

发布时间:2026/9/16 16:02:31 来源:尧图企业网站定制
简介一套基于PyTorch的CNN鸟类品种识别小程序源码包面向希望快速上手深度学习图像分类与前后端联调的开发者。项目通过三个Python脚本串起完整链路01生成图片路径与标签的txt文件并划分训练集和验证集02执行模型训练03启动Flask服务端提供推理接口微信小程序端通过接口展示识别结果适合学习从数据处理到部署调用的小型实战项目。压缩包共26个文件主要包括3个核心py脚本、小程序前端代码json/js/wxss/wxml、低保真示例提示图、requirements.txt依赖清单以及docx说明文档整体仅321KB。已有106人学习浏览。代码每一行均配中文注释说明文档附环境安装指引数据集文件夹可自行创建分类目录并放入对应图片灵活扩展识别种类每个类别目录内也放置了图片摆放提示图。通过这套资源可快速跑通一个包含图片收集、模型训练、服务端部署和小程序展示的完整CNN识别Demo。1. 没有图片数据集的 CNN 鸟类分类项目反而更适合练手拿到这个压缩包的第一印象是居然一个真实图片都没有但把01数据集文本生成制作.py、02深度学习模型训练.py、03flask_服务端.py三个脚本读完就会发现这个缺图反而是件好事。它把“数据准备 → 模型训练 → Flask 服务 → 微信小程序调用”整条链路拆开每个环节都需要你自己动手补全。相比那些把图片打包好、跑一下就能看到结果的 demo这种半成品更接近真实项目你得先搜集百灵鸟、画眉鸟、麻雀的图片理解目录结构再训练 CNN 卷积神经网络。适合刚学 PyTorch、准备课程设计、或者想快速搭一个图像分类原型的开发者。环境装好后按 requirement.txt 走一遍推荐用 Anaconda 建 Python 3.7 或 3.8 环境PyTorch 装 1.7.1 或 1.8.1就能得到一个能手动识别鸟类图片的小程序。2. 数据管道先行把鸟类图片变成可训练的 TXT 标签很多图像分类项目跑不出来不是模型写错而是图片路径、标签编码在第一步就乱了。这个项目不含数据集图片所以第一步反而是最关键的。下载解压后要建立一个dataset根目录下面每个子文件夹对应一种鸟例如百灵鸟、画眉鸟、麻雀。每个子文件夹里放多少张图片没有硬性要求但每个类别至少要有 20 张以上才勉强能训练类别数量可以随时增加代码里的类别判断是按子文件夹名动态生成的。2.1 数据集目录结构与 label 编码约定推荐的目录结构常见做法是dataset/ ├── 百灵鸟/ │ ├── 1.jpg │ └── 2.jpg ├── 画眉鸟/ │ └── 1.jpg └── 麻雀/ └── 1.jpg脚本遍历时会自动把“百灵鸟”“画眉鸟”“麻雀”这类文件夹名按照排序后的顺序映射成 0、1、2 的整数标签。好处是标签和文件名解耦后续增加一个“喜鹊”文件夹时只需要重新运行 01 脚本新类别会被分配一个新 id模型最后一层的num_classes同步改一下即可。压缩包里每个文件夹放了一张提示图告诉你图片该放在什么位置这是很贴心的小设计。生成的文件约定如下文件内容消费方train.txt图片绝对路径 数字标签每行一组02 训练脚本val.txt图片绝对路径 数字标签每行一组02 训练脚本用于验证类别名称由文件夹名决定不写死后续 Flask 端需手动同步2.2 提取路径和标签写入 train.txt 与 val.txt01数据集文本生成制作.py的核心逻辑并不复杂用标准库就能完成。它处理的是“路径 标签”两条信息而不是把图片直接读进内存。这么做可以减少训练前的内存占用也为后续 DataLoader 按需读取图片做准备。import os import random def write_txt(path, lines): with open(path, w, encodingutf-8) as f: for item in lines: f.write(f{item[0]} {item[1]}\n) def generate_text(root_dir, train_ratio0.8): categories sorted(os.listdir(root_dir)) label_map {name: idx for idx, name in enumerate(categories)} samples [] for name in categories: cat_dir os.path.join(root_dir, name) if not os.path.isdir(cat_dir): continue for file in os.listdir(cat_dir): if not file.lower().endswith((.jpg, .jpeg, .png)): continue image_path os.path.abspath(os.path.join(cat_dir, file)) samples.append((image_path, label_map[name])) random.seed(42) random.shuffle(samples) split_idx int(len(samples) * train_ratio) write_txt(train.txt, samples[:split_idx]) write_txt(val.txt, samples[split_idx:]) print(f共生成 {len(samples)} 张图片 f训练集 {split_idx} 张验证集 {len(samples)-split_idx} 张)这段代码有两个容易踩的坑。第一个是random.seed(42)必须放在shuffle之前否则每次划分结果不一致后面调参时无法对比准确率。第二个是os.path.abspath会把相对路径转成绝对路径因为训练脚本可能从不同目录启动用相对路径很容易出现“找不到图片”的 FileNotFoundError。train_ratio默认 0.8在图片总数不足 50 张时可以改成 0.9验证集只留少量做参考。生成出来的 TXT 每行是“绝对路径 数字标签”中间用空格分隔。训练脚本 DataLoader 拿到这些路径后再配合PIL.Image.open读取。如果你要在数据集上做样本均衡可以在这里统计每个类别的样本数并打印到控制台方便决定是否做过采样。3. CNN 训练从零搭建卷积神经网络并理解每个参数02深度学习模型训练.py做的是真正训练模型的工作。它读取上一步生成的 TXT用 PyTorch 的 DataLoader 加载图片把 RGB 图片统一缩放成 3×224×224 的张量再输入到自建的 CNN 卷积神经网络里。为什么不直接用 ResNet 预训练模型因为这个项目定位是课程设计和入门 demo去掉预训练权重能让你看清卷积、池化、全连接层各自的作用而且三层卷积网络在 CPU 上也能在几分钟内跑完。3.1 网络结构三层卷积加全连接以 3 分类为例网络结构设计如下import torch.nn as nn class BirdCNN(nn.Module): def __init__(self, num_classes3): super(BirdCNN, self).__init__() self.conv_layers nn.Sequential( nn.Conv2d(3, 32, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(32, 64, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(64, 128, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2) ) self.classifier nn.Sequential( nn.Dropout(0.5), nn.Linear(128 * 28 * 28, 256), nn.ReLU(inplaceTrue), nn.Linear(256, num_classes) ) def forward(self, x): x self.conv_layers(x) x x.view(x.size(0), -1) return self.classifier(x)输入是 224×224 的 RGB 图像经过三次 MaxPool2d(2) 后空间尺寸变成 28×28所以全连接层第一维是 128×28×28。注意如果你把输入改成其他尺寸这个数字要跟着调整。一个更稳的做法是在卷积层后用nn.AdaptiveAvgPool2d((1, 1))替代x.view这样不管输入多少尺寸全连接层输入维度都是 128。这个项目为了直观保留了固定尺寸计算是初学者最容易改错的地方。3.2 DataLoader 与数据增强训练脚本里通常会定义这样的数据变换与加载逻辑from torch.utils.data import Dataset, DataLoader from PIL import Image from torchvision import transforms transform transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) class BirdDataset(Dataset): def __init__(self, txt, transformNone): self.samples [] with open(txt, r, encodingutf-8) as f: for line in f: path, label line.strip().split() self.samples.append((path, int(label))) self.transform transform def __len__(self): return len(self.samples) def __getitem__(self, idx): path, label self.samples[idx] image Image.open(path).convert(RGB) if self.transform: image self.transform(image) return image, label train_loader DataLoader( BirdDataset(train.txt, transform), batch_size32, shuffleTrue, num_workers0 )这里RandomHorizontalFlip是常见的图像增强 CNN 算法之一对鸟类姿态变化有正面作用。num_workers在 Windows 下建议设 0否则容易在 DataLoader 里报多进程相关错误Linux 下可以适当提高。训练主循环不复杂每个 epoch 遍历一次 train_loader用交叉熵损失和 Adam 优化器更新权重。3.3 训练主循环与模型保存import torch import torch.nn as nn device torch.device(cuda if torch.cuda.is_available() else cpu) model BirdCNN(num_classes3).to(device) criterion nn.CrossEntropyLoss() optimizer torch.optim.Adam(model.parameters(), lr0.001) best_acc 0.0 for epoch in range(30): model.train() running_loss 0.0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() val_acc evaluate(model, val_loader, device) print(fEpoch {epoch1}/30, Loss: {running_loss:.4f}, Val Acc: {val_acc:.4f}) if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), best_model.pth)evaluate是在验证集上逐批前向计算准确率的函数记得在函数里用model.eval()和torch.no_grad()包裹否则推理图会一直累积内存。torch.save建议只保存state_dict不保存整个 model 对象因为 Flask 端重新实例化后加载 weights 更安全。3.4 超参数推荐超参数推荐值说明输入尺寸224×224与网络全连接层尺寸绑定卷积核3×3参数少适合小数据集池化2×2 MaxPool下采样降低特征分辨率优化器Adam, lr0.001收敛稳定不必手动调 lrbatch_size16/32显存不足时先用 16epochs2030观察验证集准确率是否平稳以上超参数适合几百张小图的数据集。如果自己收集的图片数量很少最好再加随机旋转、亮度抖动这类增强避免模型过拟合到背景纹理上。4. Flask 服务端把训练好的模型包装成可调用的 HTTP 接口03flask_服务端.py是整个链路的“中间层”。训练好的模型是一个.pth文件小程序不能直接加载 PyTorch所以需要用轻量服务把模型包起来。Flask 是这里最合适的选择开发快、依赖少、能和 PyTorch 无缝配合。服务端接收到图片后完成预处理、推理、返回 JSON小程序只需要关心上传图片和显示结果。4.1 加载模型和图像预处理为了不让模型在每次请求时重新初始化常见做法是把模型加载放到全局位置只执行一次。同时要复用训练时的 transform 逻辑让测试预处理和训练预处理保持一致否则会出现训练 96% 验证 70% 的错位。import torch from flask import Flask, request, jsonify from torchvision import transforms from PIL import Image model BirdCNN(num_classes3) model.load_state_dict(torch.load(best_model.pth, map_locationcpu)) model.eval() transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) LABELS [百灵鸟, 画眉鸟, 麻雀] app Flask(__name__)这里map_locationcpu很关键因为训练时可能用了 GPU服务端部署在无显卡机器上时必须显式映射到 CPU否则会直接报RuntimeError: Attempting to deserialize object on a CUDA device。LABELS列表的顺序要和训练脚本里的 class 映射保持一致这是最容易出错的环节。4.2 接收图片并返回推理结果小程序端使用wx.uploadFile时图片以表单字段形式传输服务端通过request.files拿到文件对象。下面的接口兼容图片字段名fileapp.route(/predict, methods[POST]) def predict(): if file not in request.files: return jsonify({error: no file}), 400 file request.files[file] image Image.open(file.stream).convert(RGB) input_tensor transform(image).unsqueeze(0) with torch.no_grad(): outputs model(input_tensor) probs torch.softmax(outputs, dim1) conf, pred torch.max(probs, 1) return jsonify({ label: LABELS[pred.item()], confidence: round(conf.item(), 4) })torch.no_grad()在推理时关闭梯度计算能明显降低内存占用。softmax把 logits 转成 0 到 1 的概率便于小程序端直接渲染百分比。如果你希望接口返回前 3 个候选类别可以用torch.topk一次取 top3 概率而不是只拿最大值这对“识别失败时给用户提示”很有用。4.3 本地联调与常见错误Flask 默认只监听127.0.0.1小程序模拟器可以直接访问 localhost但真机调试时必须使用局域网地址。启动时指定host0.0.0.0if __name__ __main__: app.run(host0.0.0.0, port5000)运行后先不要急着打开小程序先用 curl 测一下接口确认服务状态再往小程序端走。常见错误如下现象原因处理方式上传后返回 400表单字段名不是file修改wx.uploadFile的name加载模型报 CUDA 错误训练用了 GPU部署机器无 GPUload_state_dict加map_locationcpu请求超时原图太大上传耗时过长小程序端先wx.compressImage分类结果全是同一类训练数据类别不均衡增补图片或做过采样重新训练5. 微信小程序端上传图片、动态标题与加载页的落地细节最后一步是让用户在微信小程序里选择鸟类图片上传到 Flask再把识别结果展示出来。压缩包里的小程序部分包含app.json、pages/index、pages/logs、pages/new0等结构重点是pages/index的上传逻辑。5.1 页面注册与小程序头部标题配置在app.json里需要把页面路径写全window字段控制小程序头部标题和导航栏样式。常见做法是在页面级配置里设置静态标题动态标题则在onLoad后通过wx.setNavigationBarTitle修改。比如识别完成后把顶部标题改成“识别结果麻雀”。{ pages: [pages/index/index, pages/new0/new0], window: { navigationBarTitleText: 鸟类识别, navigationBarBackgroundColor: #f0f0f0 } }5.2 调用 wx.uploadFile 上传图片在index.js中先用wx.chooseMedia选择图片再把临时路径通过wx.uploadFile上传到 Flask。下面是一段带注释的完整流程wx.chooseMedia({ count: 1, mediaType: [image], success(res) { const tempFilePath res.tempFiles[0].tempFilePath wx.uploadFile({ url: http://127.0.0.1:5000/predict, filePath: tempFilePath, name: file, success(uploadRes) { const data JSON.parse(uploadRes.data) wx.setNavigationBarTitle({ title: data.label data.confidence }) }, fail() { wx.showToast({ title: 上传失败 }) } }) } })这里name必须等于 Flask 端request.files里的字段名。如果开发者工具报uploadFile: fail url not in domain list说明没有关闭合法域名校验正式发布前必须把 Flask 服务放到 HTTPS 域名下并在小程序后台配置 request 合法域名。修改刚进入的加载页面时可以调整app.wxss里的全局背景色更实用的做法是在index页面onLoad里展示骨架屏识别完成后再替换结果区域。5.3 让结果页保持可复用最后留一个实用技巧如果改天把鸟类换成花卉图片分类只需要重新生成数据集文本、修改 Flask 端的LABELS和num_classes小程序端代码一行动都不用改。因为整条链路里类别名称只出现在服务端响应里小程序只渲染 label 和 confidence属于一种前后端解耦的方式。这样后续加新类别只需要在数据集目录里新建文件夹重跑 01、02、03 三个脚本再重启 Flask 就好。本文还有配套的精品资源点击获取

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价