PyTorch图像识别+Flask部署:宠物分类端到端实战
简介这是一份面向深度学习入门者与计算机视觉爱好者的宠物图像识别实战项目源码基于PyTorch构建分类模型并用Flask封装后端推理接口帮助读者理解从数据采集、模型训练到服务部署的完整链路。压缩包共约2000个文件以1993张jpg宠物图片为主要数据样本另含4个Python脚本、2个JSON配置与1个Markdown说明文档整体约34.73MB覆盖猫、犬、爬行动物、两栖动物等多类别图像。其中训练脚本负责模型构建与分类爬虫脚本用于采集网络宠物图片预测脚本支持单张或批量识别JSON文件记录训练损失与准确率以及类别定义Flask接口则便于前端调用识别服务。项目目录还区分了图像数据与训练日志结构清晰适合作为课程设计、毕业设计或自学练手参考。目前已有41人学习下载可帮助读者快速跑通一套可复用的图像识别流程。1. 从一张猫图说起PyTorch Flask 的宠物识别到底在做什么你拍了张照片想立刻知道这是布偶还是暹罗最直接的做法是本地跑一个 PyTorch 模型推理再用 Flask 把它包成一个网页接口手机浏览器打开就能传图看结果。这个标题讲的就是这条链路PyTorch 负责图像识别算法本身Flask 负责把它变成能访问的网页服务中间靠一个训练好的分类模型串起来。适合谁手上有标注好的宠物图片、想快速验证一个端到端方案的人或者已经会写 PyTorch 训练脚本、但不知道怎么让非技术同事也能用起来的人。它不解决“识别率从 85% 提到 99%”这种模型调优问题解决的是“模型跑通之后怎么让别人也能用”的落地问题。我见过太多人卡在这一步训练脚本跑得飞起一到部署就翻车要么环境对不上要么接口传参写错要么图片预处理和训练时不一致导致线上效果玄学下降。这篇就按我实际做过的路径把 PyTorch 图像识别加 Flask 部署这条线拆开讲清楚从环境搭建到接口联调再到几个必踩的坑让你能照着复现。2. 环境搭建与模型选型别在第一步就卡住2.1 PyTorch 安装CPU 还是 GPU先想清楚再动手很多人一上来就搜“pytorch安装教程超详细”结果被 CUDA 版本、显卡驱动、WSL 绕晕。我的建议很直接如果你只是做宠物图像识别这种中小规模分类任务推理阶段 CPU 完全够用训练阶段再考虑 GPU。先装 CPU 版本把链路跑通后面要提速再换别一开始就追求 GPU 环境容易在驱动版本上耗掉半天。用 conda 建一个干净环境这是最稳的做法conda create -n petcls python3.10 -y conda activate petcls # 安装 CPU 版 PyTorch版本按官网当前稳定版来 pip install torch torchvision --index-url https://download.pytorch.org/whl/cpu # Flask 和图像处理依赖 pip install flask pillow numpy逻辑说明单独建环境是为了避免和系统里其他项目的 torch 版本冲突这是血泪经验混装之后报错很难查。torchvision负责图像变换和预训练模型加载pillow处理上传的图片numpy做数组转换。参数上Python 3.10 是我目前用得最顺的版本3.11 以上有些旧版 torch 轮子不全3.8 又偏老。如果你确实要用 GPU把--index-url换成对应 CUDA 版本的源但注意先确认显卡驱动支持的最高 CUDA 版本别直接装最新。提示安装完用python -c import torch; print(torch.__version__)验证能打印出版本号才算成功不要凭感觉。2.2 模型选型ResNet18 够用别一上来就上大模型宠物图像识别本质是细粒度分类猫狗品种之间差异不大。我一般会先用 ResNet18 做基线原因是它参数量小、推理快、预训练权重容易拿在几千张图的宠物数据集上微调就能到可用的准确率。如果你数据量特别大、品种特别多再考虑 ResNet50 或 EfficientNet但部署时模型体积和推理延迟会明显上升。加载预训练模型并替换分类头的写法import torch import torch.nn as nn from torchvision import models def build_model(num_classes): # 加载 ImageNet 预训练权重迁移学习能省大量标注数据 model models.resnet18(weightsmodels.ResNet18_Weights.DEFAULT) # 替换最后的全连接层输出改为自己的类别数 in_features model.fc.in_features model.fc nn.Linear(in_features, num_classes) return model model build_model(num_classes5) # 假设识别5种宠物逻辑说明weightsDEFAULT会自动下载官方预训练权重第一次运行需要联网。替换fc层是因为原模型输出 1000 类我们要改成自己的类别数。参数num_classes必须和你的标签映射一致训练时用几类推理时就得是几类否则输出维度对不上直接报错。训练部分不是这篇重点你按常规的交叉熵损失加 Adam 优化器微调几轮即可记得保存state_dict而不是整个模型对象加载时更灵活。2.3 图片预处理训练和推理必须用同一套变换这是最容易翻车的地方。训练时你用了归一化、缩放、中心裁剪推理时如果只做Resize就直接送进模型准确率会莫名其妙掉一截。我一般把预处理定义成一个函数训练和推理共用from torchvision import transforms def get_transform(): return 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]) ])逻辑说明Resize(256)加CenterCrop(224)是 ImageNet 系列模型的标准输入尺寸Normalize的均值和方差也是官方预训练时用的必须保持一致。如果你的训练脚本里用了随机裁剪增强推理时要去掉随机部分只保留确定性的缩放和裁剪。参数上224 是 ResNet 的默认输入换成其他模型要查对应输入尺寸别硬套。3. Flask 接口开发把模型变成能传图的网页服务3.1 最小可用的上传接口从 request 到推理结果Flask 开发的核心就三件事接收图片、预处理、返回结果。先写一个最小可跑的版本from flask import Flask, request, jsonify from PIL import Image import torch import io app Flask(__name__) model build_model(num_classes5) model.load_state_dict(torch.load(pet_model.pth, map_locationcpu)) model.eval() # 切换到推理模式这行不能省 transform get_transform() labels [布偶, 暹罗, 柯基, 柴犬, 金毛] app.route(/predict, methods[POST]) def predict(): if file not in request.files: return jsonify({error: 没有上传文件}), 400 file request.files[file] img Image.open(io.BytesIO(file.read())).convert(RGB) tensor transform(img).unsqueeze(0) # 增加 batch 维度 with torch.no_grad(): # 关闭梯度省内存提速 outputs model(tensor) prob torch.softmax(outputs, dim1) conf, idx torch.max(prob, dim1) return jsonify({ label: labels[idx.item()], confidence: round(conf.item(), 4) }) if __name__ __main__: app.run(host0.0.0.0, port5000)逻辑说明request.files拿上传的文件convert(RGB)防止 PNG 带透明通道导致三通道转换报错。unsqueeze(0)是加 batch 维度模型要求输入是[N, C, H, W]。torch.no_grad()在推理时必加否则显存或内存占用会高很多。softmax把输出转成概率max取最大概率对应的类别。参数上map_locationcpu保证在无 GPU 机器上也能加载host0.0.0.0让局域网内其他设备能访问只写127.0.0.1的话手机连不上。注意model.eval()一定要在加载权重后调用它会影响 BatchNorm 和 Dropout 的行为漏掉这行线上结果会和训练时对不上。3.2 前端页面一个表单就够别过度设计Flask 绑定网页元素最朴素的方式就是表单提交不需要前后端分离也能用!DOCTYPE html html headmeta charsetutf-8title宠物识别/title/head body form action/predict methodpost enctypemultipart/form-data input typefile namefile acceptimage/* button typesubmit识别/button /form /body /html逻辑说明enctypemultipart/form-data是文件上传必须的漏掉的话后端收不到文件。namefile要和后端request.files[file]对应。这个页面直接返回 JSON如果你想在页面上显示结果可以用 JavaScript 的fetch发请求再渲染但初期用表单验证链路更快。参数上acceptimage/*只是给浏览器一个提示不限制实际上传类型后端仍要做校验。3.3 接口联调用 curl 先验证再上浏览器写完接口别急着开浏览器先用命令行验证能快速定位是后端问题还是前端问题curl -X POST -F filetest_cat.jpg http://127.0.0.1:5000/predict逻辑说明-F表示表单上传后面跟本地图片路径。如果返回{label: ..., confidence: ...}说明后端通了。如果报 400检查字段名是不是file如果报 500看 Flask 控制台的堆栈多半是图片格式或模型加载问题。这一步能省掉大量“到底是前端没传对还是后端没接住”的扯皮时间。4. 避坑与排查那些让我加班到凌晨的细节4.1 上传大图导致内存暴涨甚至服务卡死现象用户传了一张手机原图几 MB 甚至十几 MB服务响应变慢并发几个请求后直接卡死。原因Image.open会把整张图解码进内存大图解码后占用的内存远大于文件本身再加上模型推理的中间张量内存很快吃满。解决在预处理前限制图片尺寸比如先做一次缩略img Image.open(io.BytesIO(file.read())).convert(RGB) img.thumbnail((512, 512)) # 限制最长边降低内存占用thumbnail会原地修改并保持比例比resize更省事。另外可以在 Flask 配置里限制最大上传体积app.config[MAX_CONTENT_LENGTH] 5 * 1024 * 1024超过直接拒绝。4.2 训练用 GPU 推理用 CPU 导致加载报错现象训练时保存的模型在 CPU 机器上load_state_dict报错提示找不到 CUDA 设备。原因保存时张量带有 GPU 设备信息加载时默认按原设备找。解决加载时加map_locationcpu前面代码里已经写了。如果保存的是整个模型对象而不是state_dict问题更多所以我一律建议只存state_dict。4.3 类别顺序不一致导致结果张冠李戴现象模型明明训练准确率很高线上识别结果却总是错位把布偶认成暹罗。原因训练时标签映射是{布偶: 0, 暹罗: 1}推理时labels列表顺序写反了。解决把标签映射存成 JSON 文件训练和推理都从同一个文件读别手写列表。这个坑很隐蔽因为模型输出本身没错错的是你解读输出的方式。4.4 Flask 默认单线程阻塞并发请求排队现象两个人同时上传图片第二个人要等第一个人识别完才响应。原因Flask 开发服务器默认单线程。解决开发阶段可以开threadedTrue生产环境用 gunicorn 加多 workergunicorn -w 4 -b 0.0.0.0:5000 app:app-w 4是 4 个 worker 进程按 CPU 核数调整。注意模型在每个 worker 里都会加载一份内存要留够。4.5 图片 EXIF 方向导致识别异常现象手机拍的竖图上传后识别结果很差横过来看就正常。原因手机照片带 EXIF 旋转信息PIL 默认不自动旋转模型看到的是转过的图。解决用ImageOps.exif_transpose自动纠正方向from PIL import ImageOps img ImageOps.exif_transpose(img)这行加在convert(RGB)之前能解决大部分手机图方向问题。5. 进阶技巧让这个方案真正能交付5.1 用 ONNX 导出提速顺便摆脱 PyTorch 依赖如果部署机器装 PyTorch 太重可以把模型导出成 ONNX推理用 onnxruntime体积小、启动快。导出脚本import torch model.eval() dummy torch.randn(1, 3, 224, 224) torch.onnx.export( model, dummy, pet_model.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch}, output: {0: batch}} )逻辑说明dummy是模拟输入用来追踪计算图。dynamic_axes让 batch 维度可变这样一次可以处理多张图。导出后用onnxruntime.InferenceSession加载推理代码要相应调整输入输出都变成 numpy 数组。参数上opset 版本不指定会用默认遇到不支持的算子再手动调。5.2 加一个健康检查接口方便排查服务状态上线后最怕不知道服务活着没。加一个简单接口app.route(/health) def health(): return jsonify({status: ok, model_loaded: model is not None})逻辑说明这个接口不涉及推理响应极快适合给监控系统轮询。model_loaded能帮你确认模型是否加载成功比只看进程在不在更有意义。5.3 批量推理一次请求处理多张图用户可能一次传多张逐张推理效率低。把输入拼成 batchtensors torch.stack([transform(Image.open(io.BytesIO(f.read())).convert(RGB)) for f in files]) with torch.no_grad(): outputs model(tensors) probs torch.softmax(outputs, dim1)逻辑说明torch.stack把多张图的张量叠成[N, C, H, W]模型一次前向就能出所有结果。注意显存或内存会随 N 线性增长N 太大要分批。参数上建议单次不超过 8 张再大就分块处理。5.4 一个我常用的验证习惯每次改完预处理或模型我不会直接上浏览器点而是固定用同一张测试图跑curl对比输出的 label 和 confidence 有没有突变。如果 confidence 从 0.95 掉到 0.6多半是预处理或标签顺序动了。这个习惯帮我省了很多“感觉不对但说不清哪里不对”的时间。做这类端到端方案最怕的就是链路太长、每步都差一点最后结果玄学。固定一张基准图每次改动后跑一遍是最便宜的后悔药。希望帮到你。本文还有配套的精品资源点击获取