基于ONNX的垃圾分类识别系统:从模型导出到Flask部署实战

发布时间:2026/10/10 11:10:51
基于ONNX的垃圾分类识别系统:从模型导出到Flask部署实战
简介这份资源是一套基于深度学习的垃圾分类系统实现面向具备Python基础、希望上手图像分类项目或课程设计的学习者。项目以卷积神经网络为核心通过训练垃圾图像数据实现自动识别与分类并借助ONNX格式完成模型导入便于跨框架、跨平台部署适合想了解模型导出与推理流程的开发者参考。压缩包共6个文件约12KB包含3个csv数据文件、2个py源码文件及1个pyc编译文件csv用于存放标签、用户与历史记录等数据py文件承担主程序与模型逻辑整体结构轻量便于快速阅读与二次修改。目前已有117人学习下载。读者可从中获取数据预处理、模型训练与评估、ONNX导入及分类推理的完整代码脉络理解从图像输入到类别输出的实现思路并据此搭建自己的垃圾分类演示系统或扩展为更复杂的视觉识别项目。1. 从一份 ONNX 模型文件说起垃圾分类系统到底交付了什么你拿到一个压缩包解压后看到app.py、rubbish.py、static、views、label.csv、history.csv、user_pwd.csv还有一个__pycache__。第一反应可能是模型在哪权重呢其实这个项目的核心思路很明确——训练阶段用 Python 深度学习框架完成推理阶段把模型导出成 ONNX 格式再由 Web 服务加载 ONNX 做前向计算。这样做的好处是部署侧不再依赖训练框架ONNX Runtime 的安装体积和启动速度都更可控适合在边缘设备或普通云主机上跑。这套系统解决的是「拍一张垃圾照片自动告诉你属于哪一类」的问题。适合两类人一是想找一个能跑通的深度学习项目做课程设计或毕业设计的同学二是需要把图像分类模型快速封装成 Web 接口的工程师。它不追求 SOTA 精度但胜在结构完整——有登录、有历史记录、有标签映射、有前端页面是一个能直接演示的闭环。2. 拆开压缩包目录结构与 ONNX 推理链路2.1 每个文件在系统里扮演什么角色先把目录摊开看。app.py是 Flask 入口负责注册路由、启动服务rubbish.py大概率封装了模型加载和推理逻辑views目录放的是蓝图或路由处理函数static存前端静态资源label.csv是类别索引到中文名称的映射表history.csv记录每次识别的结果user_pwd.csv存用户凭证。__pycache__是 Python 字节码缓存可以忽略。文件/目录作用是否可替换app.pyFlask 启动入口注册蓝图否rubbish.pyONNX 模型加载与推理封装可替换模型路径views/路由处理逻辑否static/CSS/JS/图片可替换label.csv类别索引→中文标签需与模型输出对齐history.csv识别历史记录可清空user_pwd.csv用户账号密码建议改为数据库这里最关键的其实是label.csv和 ONNX 模型的输出维度必须对齐。常见做法是模型输出 4 类或 6 类label.csv就对应 4 行或 6 行顺序不能错。一旦顺序错了模型明明预测的是「可回收物」页面上却显示「厨余垃圾」这种翻车在现场演示时非常尴尬。2.2 ONNX 模型加载与推理的最小代码骨架下面这段代码是我根据这类项目的常见写法还原的推理骨架放在rubbish.py里。它用onnxruntime加载模型用PIL做预处理最后返回类别索引和置信度。import onnxruntime as ort import numpy as np from PIL import Image # 加载 ONNX 模型指定 CPU 执行提供者 session ort.InferenceSession(model.onnx, providers[CPUExecutionProvider]) # 获取输入名称和形状通常是 [1, 3, 224, 224] input_name session.get_inputs()[0].name input_shape session.get_inputs()[0].shape def preprocess(image_path): img Image.open(image_path).convert(RGB) img img.resize((224, 224)) # 与训练时输入尺寸一致 arr np.array(img).astype(np.float32) / 255.0 # 归一化到 [0,1] mean np.array([0.485, 0.456, 0.406]) std np.array([0.229, 0.224, 0.225]) arr (arr - mean) / std # ImageNet 标准化 arr arr.transpose(2, 0, 1) # HWC - CHW arr np.expand_dims(arr, axis0) # 增加 batch 维度 return arr def predict(image_path): tensor preprocess(image_path) outputs session.run(None, {input_name: tensor}) logits outputs[0][0] idx int(np.argmax(logits)) confidence float(np.exp(logits[idx]) / np.sum(np.exp(logits))) return idx, confidence逻辑说明ort.InferenceSession是 ONNX Runtime 的标准入口providers参数决定用 CPU 还是 GPU。预处理里的均值和标准差必须和训练时一致否则精度会掉。session.run的第一个参数传None表示返回所有输出第二个参数是输入字典。最后用 softmax 把 logits 转成概率取最大值对应的索引。参数怎么改如果模型输入是 320x320就把resize改掉如果训练时没有做 ImageNet 标准化就把 mean/std 那两行去掉如果模型输出已经是 softmax 后的概率就不需要再算 exp。2.3 Flask 路由如何把推理结果送到前端app.py里通常会有一个/upload或/predict路由接收前端上传的图片调用predict再把结果渲染回页面。下面是一个最小可用的路由写法from flask import Flask, request, render_template import os from rubbish import predict app Flask(__name__) UPLOAD_FOLDER static/uploads os.makedirs(UPLOAD_FOLDER, exist_okTrue) app.route(/predict, methods[POST]) def do_predict(): file request.files[image] save_path os.path.join(UPLOAD_FOLDER, file.filename) file.save(save_path) idx, conf predict(save_path) # 读取 label.csv 做索引映射 with open(label.csv, r, encodingutf-8) as f: labels [line.strip() for line in f.readlines()] label labels[idx] return render_template(result.html, labellabel, confidenceround(conf, 4))这段代码的关键点是上传目录必须存在否则file.save会直接抛异常label.csv的读取顺序要和模型输出索引严格对应confidence传给前端时最好保留四位小数避免页面上出现一长串浮点数。3. 把 PyTorch 模型转成 ONNX导出参数与验证方法3.1 导出时的三个核心参数如果你手上有训练好的 PyTorch 权重想替换掉项目自带的 ONNX 模型导出这一步绕不开。torch.onnx.export有三个参数最容易出问题opset_version、input_names、dynamic_axes。import torch import torchvision.models as models model models.resnet18(pretrainedFalse) model.fc torch.nn.Linear(512, 4) # 假设 4 分类 model.load_state_dict(torch.load(best.pth, map_locationcpu)) model.eval() dummy_input torch.randn(1, 3, 224, 224) torch.onnx.export( model, dummy_input, model.onnx, opset_version11, # 常用 11 或 12太低不支持某些算子 input_names[input], # 与推理代码里的 input_name 对应 output_names[output], dynamic_axes{input: {0: batch}, output: {0: batch}} )opset_version选 11 是比较稳妥的ONNX Runtime 对 11 的支持最成熟。dynamic_axes允许 batch 维度动态变化这样推理时可以一次传多张图。如果导出时报「Unsupported operator」先升级opset_version再检查模型里有没有自定义层。3.2 导出后怎么验证 ONNX 和原模型输出一致导出完成不代表万事大吉。我一般会做一次数值对齐用同一张输入图片分别跑 PyTorch 和 ONNX Runtime比较两者的输出差异。import onnxruntime as ort import numpy as np import torch # PyTorch 输出 with torch.no_grad(): pt_out model(dummy_input).numpy() # ONNX 输出 sess ort.InferenceSession(model.onnx) onnx_out sess.run(None, {input: dummy_input.numpy()})[0] # 比较最大绝对误差 diff np.max(np.abs(pt_out - onnx_out)) print(最大误差:, diff) # 一般应小于 1e-4如果误差超过 1e-3说明导出过程中有算子被近似替换了常见于AdaptiveAvgPool或自定义激活函数。这时候要么换 opset要么把模型结构改得更「标准」一些。3.3 用 Netron 看一眼模型结构导出后建议用 Netron 打开.onnx文件确认输入输出名称、维度、算子类型。这一步能提前发现很多问题比如输入名称不是input或者输出维度是[1, 1000]而不是你期望的[1, 4]。Netron 是图形化工具不需要写代码拖进去就能看。4. 避坑与排查从环境到标签的五个血泪经验4.1 现象启动 Flask 报ModuleNotFoundError: No module named onnxruntime原因环境里没装 ONNX Runtime或者装的是 GPU 版但机器没有 CUDA。解决pip install onnxruntime装 CPU 版即可除非你确认要上 GPU。如果已经装了onnxruntime-gpu但报错先卸载再装 CPU 版。4.2 现象上传图片后页面显示「Internal Server Error」原因大概率是label.csv的编码问题。Windows 下用 Excel 编辑过 CSV保存时可能变成 GBK 编码而代码里用utf-8读取就会崩。解决用 VS Code 或 Notepad 把label.csv转成 UTF-8 无 BOM 格式或者代码里加encodingutf-8-sig。4.3 现象预测结果永远是同一个类别原因预处理没做对。常见情况是训练时用了归一化推理时忘了或者输入通道顺序搞错把 RGB 当成 BGR。解决对照训练脚本里的transforms.Normalize参数逐行核对推理预处理。另外检查label.csv行数是否和模型输出维度一致。4.4 现象ONNX 模型加载成功但推理速度很慢原因ONNX Runtime 默认用 CPU 单线程或者模型输入尺寸太大。解决在InferenceSession里加sess_options.intra_op_num_threads 4或者把输入从 448 降到 224。如果机器有 GPU可以装onnxruntime-gpu并指定CUDAExecutionProvider。4.5 现象history.csv越写越大页面加载变慢原因每次识别都追加一行没有清理机制。解决定期归档或只保留最近 1000 条。也可以在写入时加一个判断超过阈值就重写文件。这个坑在演示阶段不明显但跑几天就能感觉到。5. 进阶技巧用 ONNX Runtime 做批量推理与量化5.1 批量推理一次处理多张图片单张推理在演示时够用但如果要处理一个文件夹的图片逐张调用session.run效率很低。ONNX 模型如果导出了动态 batch 维度就可以一次传多张。def batch_predict(image_paths): tensors [preprocess(p) for p in image_paths] batch np.concatenate(tensors, axis0) # [N, 3, 224, 224] outputs session.run(None, {input_name: batch})[0] indices np.argmax(outputs, axis1) return indices.tolist()这里的关键是np.concatenate把多张图的张量拼成一个 batch。注意显存或内存占用会随 batch 增大而线性增长一般设 batch8 或 16 比较稳。5.2 INT8 量化把模型体积压到四分之一ONNX Runtime 提供了训练后量化工具可以把 FP32 模型转成 INT8体积缩小约 4 倍推理速度也能提升。代价是精度可能掉 1~3 个百分点。from onnxruntime.quantization import quantize_dynamic, QuantType quantize_dynamic( model_inputmodel.onnx, model_outputmodel_int8.onnx, weight_typeQuantType.QUInt8 )量化后的模型直接用ort.InferenceSession(model_int8.onnx)加载即可代码不用改。如果发现精度掉得太多可以改用quantize_static并提供一个校准数据集但配置会复杂一些。5.3 一个我踩过的坑量化后标签错乱有一次我量化完模型发现预测结果全乱了。排查半天才意识到量化脚本默认会优化模型结构某些情况下会改变输出节点的顺序。解决办法是量化后重新用 Netron 确认输出维度并在推理代码里打印一次outputs[0].shape。从那以后我每次量化完都强制走一遍数值对齐确认最大误差在可接受范围内才上线。希望帮到你。本文还有配套的精品资源点击获取