ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

PyTorch+Flask宠物图像识别:从训练到部署全链路实战

PyTorch+Flask宠物图像识别:从训练到部署全链路实战 简介这份资源是基于PyTorch与Flask构建的宠物图像识别完整项目包面向具备一定深度学习基础、希望打通从模型训练到Web服务部署全流程的开发者与学习者。包内共2000个文件以1993张jpg宠物图片作为训练与测试样本辅以4个Python脚本、2个JSON配置与1个Markdown文档压缩包约34.73MB。其中classify.py负责模型分类predict.py支持单张或批量预测api.py基于Flask对外提供识别接口crawling.py承担网络图片爬取train_loss_accuracy.json与classes.json分别记录训练指标和类别定义readme.md则给出部署与接口说明。图片覆盖猫、犬、爬行动物、两栖动物等多类宠物目录结构清晰便于直接复现训练与推理。目前已有41人学习下载适合作为课程设计、毕业项目或图像识别入门实战的参考方案。1. 宠物图像识别从零落地PyTorch 训练 Flask 上线的完整链路家里猫狗生病那会儿我拍了几十张照片想快速判断品种和可能的健康风险翻遍手机相册才发现——手动分类根本扛不住。这就是「基于 PyTorch 和 Flask 的宠物图像识别」要解决的问题用 PyTorch 训练一个能分辨猫、狗等常见宠物的图像分类模型再用 Flask 把它包成一个网页端接口上传图片就能拿到识别结果。整套方案适合两类人一是刚学完 PyTorch 基础、想找一个完整项目练手的开发者二是需要把图像识别能力嵌进自己 Web 应用、但不想从零造轮子的工程师。它不追求 SOTA 精度追求的是「训练→导出→部署→调通」这条链路能跑通、能复现、能改。下面按环境搭建、数据准备、模型训练、Flask 服务、避坑排查、进阶优化六步走每一步都给出可抄的命令和参数。2. 环境搭建与依赖选型PyTorch 装 CPU 还是 GPU 版2.1 先确认硬件和系统再决定装哪个版本装 PyTorch 之前先跑一条命令看清楚自己有什么# Linux / WSL 下查看显卡和驱动 nvidia-smi # 查看 Python 版本建议 3.9 ~ 3.11 python --version # 查看 CUDA 版本如果有 NVIDIA 显卡 nvcc --versionnvidia-smi右上角显示的CUDA Version是驱动支持的最高 CUDA 版本不是你当前安装的版本。PyTorch 官网的安装命令里cu121、cu118指的是 PyTorch 编译时链接的 CUDA 版本只要不超过驱动支持的上限就能跑。没有 NVIDIA 显卡、或者用的是 AMD 显卡比如 7900XTX 在 WSL 下跑 PyTorch 需要走 ROCm 路线直接装 CPU 版最省事训练小规模宠物数据集完全够用。我一般会新建一个 conda 环境避免和系统 Python 打架conda create -n petcls python3.10 -y conda activate petcls # CPU 版 PyTorch体积小、装得快 pip install torch torchvision --index-url https://download.pytorch.org/whl/cpu # 如果有 CUDA 12.1换成下面这行 # pip install torch torchvision --index-url https://download.pytorch.org/whl/cu121参数说明--index-url指定 PyTorch 官方 wheel 源比默认 PyPI 快很多torchvision里自带 ResNet、MobileNet 等预训练模型和图像变换工具是本项目的主力依赖。装完用下面这段验证import torch print(torch.__version__) print(CUDA available:, torch.cuda.is_available()) # 有 GPU 时打印设备名没有则输出 cpu print(Device:, torch.cuda.get_device_name(0) if torch.cuda.is_available() else cpu)如果torch.cuda.is_available()返回 False 但你确实有 NVIDIA 显卡九成是装成了 CPU 版或者驱动版本低于 PyTorch 要求的 CUDA 版本。这时候别急着重装系统先pip uninstall torch torchvision再按官网对应命令重装即可。2.2 Flask 和辅助库的安装Flask 本身很轻但图像上传、推理、跨域这几件事需要几个搭档pip install flask flask-cors pillow numpyflaskWeb 框架本体负责路由和请求处理flask-cors前后端分离时解决跨域如果 Flask 直接渲染 HTML 可以不装pillow读图、缩放、转 RGB推理前的预处理全靠它numpy把 PIL 图像转成模型能吃的数组版本上不用太纠结Flask 2.x 和 3.x 的 API 在本项目用到的部分基本一致。真正容易翻车的是 Pillow 版本和 torchvision 的 transforms 兼容性——如果遇到AttributeError: module PIL.Image has no attribute ANTIALIAS说明 Pillow 升到 10 以上而 torchvision 还是老版本降 Pillow 到 9.5 或升 torchvision 都能解决。3. 数据准备与模型训练用迁移学习把 ResNet 改成宠物分类器3.1 数据集目录结构和预处理参数宠物图像识别最常见的数据组织方式是ImageFolder要求的按类别分文件夹dataset/ ├── train/ │ ├── cat/ │ │ ├── 001.jpg │ │ └── ... │ ├── dog/ │ └── rabbit/ └── val/ ├── cat/ ├── dog/ └── rabbit/每个类别一个文件夹文件夹名就是标签。训练集和验证集按 8:2 或 7:3 切分类别数量根据你的实际需求定猫狗二分类是最小可用版本。图像尺寸统一缩到 224×224这是 ResNet 系列的标准输入。预处理用 torchvision 的 transforms 组合from torchvision import transforms train_tf transforms.Compose([ transforms.Resize((224, 224)), # 统一尺寸 transforms.RandomHorizontalFlip(), # 随机水平翻转增广 transforms.RandomRotation(15), # 随机旋转 ±15 度 transforms.ToTensor(), # 转成 [0,1] 的张量 transforms.Normalize( # 按 ImageNet 均值方差归一化 mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225] ), ]) val_tf transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize( mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225] ), ])Normalize里的均值和方差是 ImageNet 统计出来的用预训练权重就必须用这套参数否则输入分布对不上精度会掉得莫名其妙。训练集加随机翻转和旋转是为了让模型见过更多姿态验证集不能加否则评估结果不稳定。3.2 迁移学习改造 ResNet18 并训练宠物数据集通常几千到几万张从零训练容易过拟合迁移学习是标准做法加载 ImageNet 预训练的 ResNet18把最后的全连接层换成你自己的类别数。import torch import torch.nn as nn from torchvision import models, datasets from torch.utils.data import DataLoader # 1. 加载预训练 ResNet18 model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) # 2. 冻结前面的卷积层只训练最后的分类头 for param in model.parameters(): param.requires_grad False # 3. 替换全连接层输出类别数按你的数据集改 num_classes 3 # cat / dog / rabbit model.fc nn.Linear(model.fc.in_features, num_classes) # 4. 数据加载 train_ds datasets.ImageFolder(dataset/train, transformtrain_tf) val_ds datasets.ImageFolder(dataset/val, transformval_tf) train_loader DataLoader(train_ds, batch_size32, shuffleTrue, num_workers4) val_loader DataLoader(val_ds, batch_size32, shuffleFalse, num_workers4) # 5. 训练配置 device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device) criterion nn.CrossEntropyLoss() # 只优化 fc 层的参数 optimizer torch.optim.Adam(model.fc.parameters(), lr1e-3) # 6. 训练循环 for epoch in range(10): model.train() running_loss 0.0 for imgs, labels in train_loader: imgs, labels imgs.to(device), labels.to(device) optimizer.zero_grad() outputs model(imgs) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() print(fEpoch {epoch1}, loss{running_loss/len(train_loader):.4f}) # 7. 保存权重 torch.save(model.state_dict(), pet_resnet18.pth)关键参数说明batch_size32在 8GB 显存下跑 ResNet18 比较稳显存小就降到 16 或 8lr1e-3是只训练分类头时的常用学习率如果解冻更多层要降到 1e-4 级别num_workers4在 Windows 下有时会报错改成 0 即可。冻结卷积层只训练 fc10 个 epoch 通常就能到 90% 以上的验证精度想再高就解冻最后两个 block 做微调但学习率要调小。3.3 验证集评估和模型导出训练完必须看验证集表现不能只看训练 lossmodel.eval() correct, total 0, 0 with torch.no_grad(): for imgs, labels in val_loader: imgs, labels imgs.to(device), labels.to(device) outputs model(imgs) _, preds torch.max(outputs, 1) correct (preds labels).sum().item() total labels.size(0) print(fVal Acc: {correct/total:.4f})如果验证精度远低于训练精度说明过拟合加数据增广或加 Dropout如果两者都低说明欠拟合解冻更多层或增加 epoch。导出时只存state_dict比存整个模型更灵活加载时先实例化结构再load_state_dict避免 pickle 反序列化带来的版本兼容问题。4. Flask 服务封装把模型变成上传图片就能调的接口4.1 最小可用的推理接口Flask 这边要做三件事接收上传的图片、预处理成模型输入、返回预测类别和置信度。from flask import Flask, request, jsonify from PIL import Image import torch import torch.nn as nn from torchvision import models, transforms import io app Flask(__name__) # 与训练时完全一致的预处理 infer_tf transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize( mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225] ), ]) # 加载模型结构 权重 CLASSES [cat, dog, rabbit] device torch.device(cuda if torch.cuda.is_available() else cpu) model models.resnet18(weightsNone) model.fc nn.Linear(model.fc.in_features, len(CLASSES)) model.load_state_dict(torch.load(pet_resnet18.pth, map_locationdevice)) model model.to(device) model.eval() app.route(/predict, methods[POST]) def predict(): if file not in request.files: return jsonify({error: no file}), 400 file request.files[file] img Image.open(io.BytesIO(file.read())).convert(RGB) tensor infer_tf(img).unsqueeze(0).to(device) with torch.no_grad(): logits model(tensor) probs torch.softmax(logits, dim1)[0] idx int(torch.argmax(probs)) return jsonify({ class: CLASSES[idx], confidence: round(float(probs[idx]), 4) }) if __name__ __main__: app.run(host0.0.0.0, port5000)逻辑说明request.files[file]对应前端表单里namefile的字段Image.open(...).convert(RGB)处理 PNG 带透明通道或灰度图的情况不转 RGB 会在ToTensor时维度对不上unsqueeze(0)是给单张图补上 batch 维度模型要求输入是[N, C, H, W]。map_locationdevice保证在 CPU 上也能加载 GPU 训练的权重。4.2 前端上传页面和联调Flask 可以直接渲染一个简单页面也可以纯做 API 让前端调。最小验证用 HTML 表单!DOCTYPE html html body form action/predict methodpost enctypemultipart/form-data input typefile namefile acceptimage/* button typesubmit识别/button /form /body /html放到templates/index.html再加一条路由app.route(/)返回render_template(index.html)即可。前后端分离时用fetch发FormData记得开flask-cors的CORS(app)否则浏览器会拦跨域请求。联调时先用curl验证接口本身通不通curl -X POST -F filetest_cat.jpg http://127.0.0.1:5000/predict返回{class:cat,confidence:0.97}就说明链路没问题。如果返回 400检查字段名是不是file如果返回 500看 Flask 控制台的堆栈多半是图片格式或模型加载路径的问题。5. 避坑与排查宠物识别项目最容易翻车的五个地方5.1 现象训练精度很高上线后识别全错原因训练时的预处理和推理时的预处理不一致。最常见的是训练用了Normalize推理时忘了加或者 Resize 尺寸对不上。解决把预处理定义抽成一个公共模块训练和推理都 import 同一份别两边各写一套。上线前用同一张图分别跑训练脚本的验证流程和 Flask 接口结果应该完全一致。5.2 现象Flask 启动报Address already in use原因5000 端口被占macOS 上还可能是 AirPlay 占用了 5000。解决换端口app.run(port5001)或者lsof -i:5000找到进程 kill 掉。生产环境别用app.run用 gunicorngunicorn -w 2 -b 0.0.0.0:5000 app:app。5.3 现象上传大图后接口超时或内存暴涨原因手机拍的照片动辄 4000×3000直接进Resize虽然能处理但解码阶段就吃了几百 MB 内存。解决在Image.open之后先做一次限制img.thumbnail((1024, 1024))再进 transforms。另外 Flask 配置MAX_CONTENT_LENGTH 5 * 1024 * 1024限制上传大小超了直接返回 413。5.4 现象Windows 下num_workers0报BrokenPipeError原因Windows 的 DataLoader 多进程实现和 Linux 不同在脚本没加if __name__ __main__保护时会递归创建进程。解决训练代码包进if __name__ __main__:或者直接把num_workers设为 0。数据量不大时num_workers0对训练速度影响有限。5.5 现象模型对某些品种完全认不出原因训练集类别不均衡或者某些品种样本太少模型偏向多数类。解决先统计每个类别的图片数量差距超过 5 倍就要处理。简单做法是WeightedRandomSampler给少样本类别更高采样权重或者在 loss 里加weight参数。更彻底的是补数据每个类别至少 200 张起步。6. 进阶技巧用 ONNX 导出把推理速度再压一截Flask PyTorch 原生推理在 CPU 上单张图大概 50~100ms如果并发上来会顶不住。把模型导出成 ONNX再用 onnxruntime 推理CPU 上通常能快 2~3 倍而且部署时不用装完整的 PyTorch镜像体积从 2GB 降到 300MB 左右。导出脚本import torch import torch.nn as nn from torchvision import models CLASSES [cat, dog, rabbit] model models.resnet18(weightsNone) model.fc nn.Linear(model.fc.in_features, len(CLASSES)) model.load_state_dict(torch.load(pet_resnet18.pth, map_locationcpu)) model.eval() dummy torch.randn(1, 3, 224, 224) torch.onnx.export( model, dummy, pet_resnet18.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch}, output: {0: batch}}, opset_version12 )dynamic_axes让 batch 维度可变这样一次可以推理多张图。opset_version12兼容性比较好遇到不支持的算子可以往上调。Flask 里换成 onnxruntimeimport onnxruntime as ort import numpy as np sess ort.InferenceSession(pet_resnet18.onnx, providers[CPUExecutionProvider]) def predict_onnx(img): tensor infer_tf(img).unsqueeze(0).numpy().astype(np.float32) outputs sess.run([output], {input: tensor})[0] probs np.exp(outputs) / np.exp(outputs).sum(axis1, keepdimsTrue) idx int(probs.argmax()) return CLASSES[idx], float(probs[0][idx])注意providers参数有 GPU 且装了 onnxruntime-gpu 时换成CUDAExecutionProvider否则用 CPU 版。softmax 这里手写是因为 ONNX 导出时没把 softmax 包含进去如果导出时模型末尾带了 softmax 层就不需要再算一遍。验证 ONNX 和 PyTorch 输出是否一致用同一张图跑两个版本置信度差距在 1e-3 以内就算通过。如果差距大检查预处理是否完全一致尤其是 Normalize 的均值和方差。这套方案我从头跑过三遍最大的教训是别在预处理上偷懒训练和推理的 transforms 必须来自同一个函数。另外 Flask 开发服务器只适合本地调试真要给别人用gunicorn nginx 是标配ONNX 是性价比最高的提速手段。希望帮到你。本文还有配套的精品资源点击获取
返回列表