☰
基于ONNX的垃圾分类识别系统:从模型导出到Flask部署实战
2026/10/10 11:10:50 网站建设 项目流程

简介:这份资源是一套基于深度学习的垃圾分类系统实现,面向具备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, axis=0) # 增加 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_ok=True) @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", encoding="utf-8") as f: labels = [line.strip() for line in f.readlines()] label = labels[idx] return render_template("result.html", label=label, confidence=round(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(pretrained=False) model.fc = torch.nn.Linear(512, 4) # 假设 4 分类 model.load_state_dict(torch.load("best.pth", map_location="cpu")) model.eval() dummy_input = torch.randn(1, 3, 224, 224) torch.onnx.export( model, dummy_input, "model.onnx", opset_version=11, # 常用 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 格式,或者代码里加encoding="utf-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, axis=0) # [N, 3, 224, 224] outputs = session.run(None, {input_name: batch})[0] indices = np.argmax(outputs, axis=1) return indices.tolist()

这里的关键是np.concatenate把多张图的张量拼成一个 batch。注意显存或内存占用会随 batch 增大而线性增长,一般设 batch=8 或 16 比较稳。

5.2 INT8 量化:把模型体积压到四分之一

ONNX Runtime 提供了训练后量化工具,可以把 FP32 模型转成 INT8,体积缩小约 4 倍,推理速度也能提升。代价是精度可能掉 1~3 个百分点。

from onnxruntime.quantization import quantize_dynamic, QuantType quantize_dynamic( model_input="model.onnx", model_output="model_int8.onnx", weight_type=QuantType.QUInt8 )

量化后的模型直接用ort.InferenceSession("model_int8.onnx")加载即可,代码不用改。如果发现精度掉得太多,可以改用quantize_static并提供一个校准数据集,但配置会复杂一些。

5.3 一个我踩过的坑:量化后标签错乱

有一次我量化完模型,发现预测结果全乱了。排查半天才意识到:量化脚本默认会优化模型结构,某些情况下会改变输出节点的顺序。解决办法是量化后重新用 Netron 确认输出维度,并在推理代码里打印一次outputs[0].shape。从那以后我每次量化完都强制走一遍数值对齐,确认最大误差在可接受范围内才上线。

希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询