☰
垃圾识别分类系统Python源码解析:从训练到部署
2026/10/1 18:53:06 网站建设 项目流程

简介:压缩包提供一套基于Python的垃圾识别分类系统的完整源码,面向需要完成课程设计、毕业设计或入门图像识别的Python学习者。项目围绕垃圾分类场景,涵盖数据集说明、模型构建与训练、测试评估、图形界面识别预测等环节,可帮助理解从数据准备到模型部署的完整流程。压缩包共28个文件,大小约1.79MB,以12个py脚本为主体,包含CNN、MobileNet训练脚本、测试脚本及窗口识别程序;同时配有项目计划书、详细设计文档、风险管理报告等Word/PDF材料,以及Markdown说明与模型目录,结构清晰、便于查阅。目前已有336人学习使用,适合作为图像分类项目的参考模板。通过研读源码,读者可掌握PIL/OpenCV等图像预处理方法、TensorFlow/Keras建模思路,并对损失曲线、模型保存加载与推理调用形成直观认识。

1. 垃圾识别分类系统Python源码:它是什么、能干什么、适合谁

一套完整的Python垃圾识别分类系统源码,通常不是一个孤立的算法脚本,而是把深度学习模型、训练流程、推理接口和数据组织方式打包在一起的工程压缩包。用户拿到zip后解压,就能看到模型定义、训练入口、预测脚本、预训练权重,以及按类别分好的图片目录。这类系统的核心任务是把一张垃圾照片(塑料瓶、纸箱、电池、果皮等)自动判定到对应类别,多数实现采用图像分类路线,复杂一点的会升级为带定位框的物体检测。它能直接解决的实际问题有两个:一是课程设计和毕业设计需要一个能跑通的完整示例;二是想做智能垃圾分类硬件(室内分类桶、小区投放点、传送带分拣)的工程师,需要一个现成基线做二次开发。适合谁来用?有Python基础和一点PyTorch使用经验的人,拿到源码包后最快半天就能把预测命令跑起来;完全零基础则建议先把环境配好,再从推理脚本读起。

2. 拆解垃圾识别源码:分类模型与检测模型的选型逻辑

2.1 先看源码用了什么模型:ResNet、MobileNet还是YOLO

解压zip后,第一件事不是急着运行,而是判断这套源码走的是哪条技术路线。因为“垃圾识别分类”这四个字在深度学习里能对应两类做法:图像分类和物体检测。图像分类的经典模型是ResNet、EfficientNet、MobileNet系列,输入是整张图片,输出是一个类别标签,适合“一个物体占画面主体”的场景,比如传送带上的单个垃圾、放在桌面上的纸盒。物体检测的经典模型是YOLO、SSD、Faster-RCNN,输出是若干个“矩形框加类别”的组合,适合画面里同时出现饮料瓶、塑料袋、纸巾的复杂投放场景。

判断源码属于哪一种,看三个线索。第一,模型定义目录里如果出现yolo、detect_head、anchor、bbox这些命名,基本是检测路线;只有resnet、mobilenet、efficientnet这样的backbone命名,大概率是纯分类。第二,训练脚本里标签的格式是关键分水岭:分类的标签是单个整数索引,检测的标签是[class_id, x_center, y_center, width, height]这样的五元组(归一化之后)。第三,推理脚本的输出形态也能看出来,分类输出predict: 可回收物,检测输出则会是detect: 塑料瓶, confidence=0.92, box=(x1,y1,x2,y2)。

我接触过的垃圾识别源码包,分类路线占大多数,原因很简单:分类的数据集好凑、训练时间短、CPU也能跑推理,交付一张静态图预测的效果图,比检测模型容易得多。如果你的应用场景是“多物品混在一起需要各自框出位置”,那就得选检测模型路线的源码,后续投入的标注成本也会直线上升。

2.2 数据流:一张垃圾图片从输入到输出经历了什么

无论源码最终用的是ResNet还是YOLO,一条完整的预测链路都包含五个环节:图像读取、预处理、模型推理、后处理和标签映射。理解这条链路,比记住某一个模型的结构重要得多,因为后续所有排错都发生在这五个环节里。

图像读取环节,最常见的是用PIL或OpenCV打开图片文件。这里埋着第一个坑:PIL读进来是RGB三通道,OpenCV读进来是BGR三通道,如果训练时用PIL做了标准化,推理时换成了OpenCV,颜色通道颠倒会让预测结果明显变差。预处理环节,源码里一般写的是transforms.Resize、transforms.ToTensor、transforms.Normalize的组合,Resize的目标尺寸取决于模型输入,常见的是224x224或256x256;Normalize的均值和标准差通常沿用ImageNet的[0.485, 0.456, 0.406],因为绝大多数源码用预训练权重初始化,数值分布必须对齐。

模型推理环节,单张图会先加一个batch维度,变成[1, 3, H, W]的张量,前向传播后得到未归一化的logits。后处理环节,分类任务对logits做softmax得到每个类别的概率,再取argmax作为预测索引;检测任务则是做NMS非极大值抑制,把重叠的候选框合并。标签映射环节,把预测索引通过一份class list文件映射回可读的中文类别名,这一步出错时系统虽然不报错,但结果会张冠李戴,后面专门讲这个坑。

2.3 压缩包里每个目录是干嘛的:按什么顺序读源码

一套规范的Python深度学习项目源码,目录结构大同小异。常见的顶层结构是:models/放网络结构定义,data/或dataset/放数据集和标签映射文件,utils/放图像预处理、可视化等工具函数,checkpoints/或weights/放预训练权重,根目录下有一个train.py和predict.py,外加requirements.txt列出依赖库。有的包还会配config.yaml统一管理超参数,或者有一个run.sh把训练脚本串起来。

我建议的阅读顺序是:先打开README.md(如果存在),看作者写了什么运行说明;再打开requirements.txt,明确依赖版本范围;接着看predict.py的main函数入口,搞清楚权重路径参数怎么传、图片路径参数怎么传;最后再读models/下的网络定义。不要一上来就钻网络结构的实现细节,那对使用源码的人来说性价比最低。

需要注意,有些分享出来的zip包里面打包方式很乱,根目录嵌套着多层文件夹,路径里还可能带中文或空格。Windows下解压后直接跑python predict.py时常出现找不到模块的情况,根源往往是当前工作目录不对。解决办法是先cd到包含train.py的目录,确认pwd路径正确,再执行脚本。

提示:拿到任何源码包,先看requirements.txt的依赖清单,再决定用conda还是venv建环境。不要直接在全局Python环境里装一堆库,后面版本冲突时后悔药可不好找。

3. 用源码跑通第一张垃圾图片预测:环境配置与最小推理命令

3.1 Python环境准备:用conda还是venv,依赖装哪些版本

跑通深度学习源码,第一道坎是环境。很多下载了源码包的用户卡在import报错上,最常见的报错是ModuleNotFoundError: No module named 'torch'或No module named 'torchvision'——这不是源码问题,是依赖没装齐。先确认Python版本,深度学习工程建议Python 3.8到3.10,太新的3.12版本可能在编译某些依赖时出兼容问题;太老的3.6版本则跑不了新版本PyTorch。

环境管理工具有两个常用选择:conda和venv。我一般这么选:如果机器上已经装了Anaconda,直接用conda建独立环境最省事;如果不想装Anaconda,用Python自带的venv也完全够用。命令行如下:

# 用venv创建名为trash_env的虚拟环境 python -m venv trash_env # 激活环境(Windows) trash_env\Scripts\activate # 激活环境(Linux/macOS) source trash_env/bin/activate # 安装依赖(有requirements.txt时优先用它) pip install -r requirements.txt # 没有requirements.txt时的最小依赖组合 pip install torch torchvision pip install Pillow numpy matplotlib opencv-python

这组命令里的关键点是:venv创建的虚拟环境会把Python库隔离在项目目录下,避免污染系统全局环境。requirements.txt里锁定的版本是作者验证过的组合,优先照装;如果机器上没有GPU,可以先把requirements里带+cu后缀的torch版本手动换成CPU版:pip install torch torchvision --index-url https://download.pytorch.org/whl/cpu。这一点不在源码说明里,但实操中几乎每个人都会撞到。

依赖装完后,验证环境的命令是python -c "import torch; print(torch.__version__)",能正常打印版本号说明环境通了。如果这一行也报错,先排查pip安装过程是否有红字报错,再检查当前终端是否真的在虚拟环境内——命令行左边有(trash_env)前缀才算激活成功。

3.2 加载预训练权重并跑通单张图片预测

环境就绪后,下一步是用源码内置的权重文件跑一次单图预测。大多数分类路线的源码会提供一个predict.py,常见的用法是传入权重路径和图片路径两个参数。如果源码没有统一入口,下面是一段可以直接替代的通用推理脚本:

# predict_single.py # 适用于分类路线的垃圾识别源码,替换成自己的模型文件名即可 import json import torch import torchvision.transforms as transforms from PIL import Image # 1. 加载模型结构(以ResNet50为例,具体类名以源码models目录为准) from torchvision.models import resnet50 model = resnet50(pretrained=False) # 不加载ImageNet权重 model.fc = torch.nn.Linear(2048, 4) # 4类垃圾:可回收/厨余/有害/其他 # 2. 加载训练好的权重 checkpoint = torch.load('checkpoints/garbage_resnet50.pth', map_location='cpu') # 兼容不同保存格式:可能是state_dict,也可能是完整模型 if isinstance(checkpoint, dict) and 'state_dict' in checkpoint: model.load_state_dict(checkpoint['state_dict']) else: model.load_state_dict(checkpoint) model.eval() # 3. 图像预处理:Resize到224x224,Normalize用ImageNet的均值和标准差 transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # 4. 读图、预处理、加batch维度(模型要求4维输入) img = Image.open('test_imgs/plastic_bottle.jpg').convert('RGB') input_tensor = transform(img).unsqueeze(0) # [1, 3, 224, 224] # 5. 推理与后处理 with torch.no_grad(): logits = model(input_tensor) # [1, 4] prob = torch.softmax(logits, dim=1) # 转成概率 pred_idx = torch.argmax(prob, dim=1).item() # 取最大概率索引 # 6. 标签映射:索引 -> 中文类别 with open('data/garbage_classes.json', 'r', encoding='utf-8') as f: classes = json.load(f) print(f"预测类别: {classes[pred_idx]},置信度: {prob[0][pred_idx].item():.4f}")

这段脚本的逻辑可以拆成三段看:第一段把模型结构和参数对应上,load_state_dict要求模型实例的结构和保存权重时完全一致,如果源码里的模型是自己定义的类而不是torchvision.models.resnet50,用我这段直接替换会报key不匹配的错误,正确做法是去源码的models/目录 import 它定义的类。第二段是预处理,Resize尺寸必须以模型定义里的输入大小为唯一标准,不是所有模型都是224。第三段是输出处理,softmax把logits归一化成概率,argmax拿到类别索引,最后经标签映射变成人类可读的名称。

跑通这一步后,检验判决标准是置信度输出。正常模型的置信度通常在0.6以上;如果输出的置信度接近0.25(4分类的随机水平)或者几个类别概率几乎相等,十有八九是权重没加载成功,或者预处理参数对不上。

3.3 从单张预测改成摄像头实时分类:改哪里、加什么

单图预测通了之后,很多人下一步是接USB摄像头做实时识别。这一步其实改动不大,核心思路就是把“读图片文件”换成“读摄像头帧”,然后在循环里逐帧调用同一个推理函数。但有两个参数必须调整:一是摄像头帧率往往达不到模型推理速度,需要加队列或跳帧处理;二是输入尺寸如果直接传1920x1080的原图,预处理Resize到224时会大幅缩放,推理耗时和内存占用都可能翻倍,建议抓帧后先裁剪或缩小。

# webcam_predict.py 摄像头实时识别的最小实现框架 import cv2 import torch from torchvision import transforms from PIL import Image # 复用上面的model和transform定义,此处省略 cap = cv2.VideoCapture(0) # 0代表默认摄像头 if not cap.isOpened(): raise RuntimeError("摄像头打开失败") frame_count = 0 while True: ret, frame = cap.read() if not ret: break frame_count += 1 if frame_count % 3 != 0: # 跳帧:每3帧推理一次,降低功耗和延迟 continue # OpenCV的BGR帧转成PIL的RGB再走预处理 rgb_img = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) pil_img = Image.fromarray(rgb_img) input_tensor = transform(pil_img).unsqueeze(0) with torch.no_grad(): prob = torch.softmax(model(input_tensor), dim=1) pred_idx = torch.argmax(prob).item() label = f"{classes[pred_idx]}: {prob[0][pred_idx].item():.2f}" cv2.putText(frame, label, (10, 30), cv2.FONT_HERSHEY_SIMPLEX, 1.0, (0, 255, 0), 2) cv2.imshow('Garbage Classifier', frame) if cv2.waitKey(1) & 0xFF == ord('q'): break cap.release() cv2.destroyAllWindows()

这个改动里最容易被忽略的是颜色空间转换:OpenCV读入的帧是BGR顺序,直接转成PIL处理会和训练时的RGB数据错位。我在实际项目里见过推理准确率从0.8+掉到0.4不到,查了半天才发现是颜色通道问题。另外跳帧的间隔不能机械照搬,它取决于推理耗时——如果单帧推理需要0.5秒,跳1帧就够;如果推理只要0.05秒,跳3帧反而会让画面卡顿。参数按实际机器性能微调。

4. 重训一个自己的垃圾识别模型:数据整理与训练参数

4.1 垃圾数据集的常见来源与目录结构转换

源码包里自带的demo数据集通常很小,每类几十张图,只够跑通流程,换到真实场景就需要重训。垃圾数据集的来源主要有三个:一是公开数据集,比如华为云垃圾分类数据集、清华大学垃圾图像分类数据集,这类数据已经按类别分好目录,下载后整理成源码要求的格式就行;二是自行拍摄,手机或相机拍实物,每类至少200张起步;三是网上爬取,注意版权和图片质量,爬来的图片很多带水印和无关背景,实际效果通常不如自行拍摄。

分类模型的目录结构最标准的是ImageFolder格式:根目录下每个类别一个子文件夹,子文件夹里是图片,目录名就是类别名。结构如下:

data/garbage_classify/ ├── train/ │ ├── recyclable/ # 可回收物图片 │ │ ├── img_001.jpg │ │ └── ... │ ├── kitchen/ # 厨余垃圾图片 │ ├── hazardous/ # 有害垃圾图片 │ └── other/ # 其他垃圾图片 ├── val/ │ ├── recyclable/ │ └── ... └── garbage_classes.json # 类别索引映射文件

从原始图片到这种结构,常见做法是写一个整理脚本,按类别把散落的图片复制到对应目录,并顺手做训练集和验证集划分:

# 建目录:mkdir_data.sh mkdir -p data/garbage_classify/train/{recyclable,kitchen,hazardous,other} mkdir -p data/garbage_classify/val/{recyclable,kitchen,hazardous,other} # 把原始图片按8:2比例随机分配到train/val(用bash的随机数实现) for img in raw_images/*.jpg; do if [ $((RANDOM % 10)) -lt 8 ]; then cp "$img" data/garbage_classify/train/kitchen/ else cp "$img" data/garbage_classify/val/kitchen/ fi done

这个脚本里的随机划分比例8:2是经验值。需要注意的是,如果原始文件夹里不同类别的图片数量不均,先做类别数量统计,少的类别可以复制增强(旋转、裁剪、颜色抖动)补到接近;差别过大的类别会让训练偏向样本多的类,导致小样本类别的准确率一塌糊涂。划分时要确保同一张图片的不同增强版本不会同时出现在train和val里,不然验证集准确率会虚高。

4.2 训练参数的取值逻辑:batch size、学习率、epochs

训练脚本的超参数决定模型最终能到什么水平。源码包里train.py默认的训练参数不一定适合你的数据集,需要按数据量和硬件重新设定。

batch size的取值逻辑:显存不足就调小,数据量大可以调大。显卡是8G显存,resnet50 + 224x224输入,batch size=32通常能跑;如果报CUDA out of memory,先把batch size减到16或8。batch size过小(比如2或4)会导致Loss震荡明显,模型难收敛,解决办法是同步调低学习率。

学习率的选择跟batch size挂钩,这个经验公式很有用:batch size翻倍,学习率也翻倍。默认基础学习率0.001搭配batch size=32是常见起点;用预训练权重做微调时,学习率可以更低,1e-4到5e-4比较稳妥,因为backbone已经学到通用特征,学习率太大了会破坏已有权重。epochs的设定取决于收敛速度:我一般先设50轮,训练时观察验证集准确率曲线,如果最后10轮验证准确率还在明显上升,就加50轮;如果连续15轮不升反降,说明已经过拟合,提前终止。

下面是一段标准训练入口命令的参考:

# 训练入口,具体参数名以源码train.py的argparse定义为准 python train.py \ --data data/garbage_classify \ --model resnet50 \ --pretrained \ --batch-size 32 \ --lr 0.001 \ --epochs 50 \ --gpu 0

这段命令的参数说明:--pretrained表示加载ImageNet预训练权重再做微调,这是小数据集下最稳妥的路线;--lr 0.001配合batch size=32是ResNet系列的常见起点;--gpu 0指定第一块显卡,没有GPU就改成--gpu -1或直接去掉这个参数。训练过程中如果发现Loss在第一个epoch后就降得很慢,不一定是代码问题,可能是学习率相对批量大小偏低了。

4.3 训练日志怎么看:损失不降、准确率虚高怎么处理

训练开始后,要盯的关键指标不是训练集准确率,而是验证集准确率和Loss曲线。这中间有很强的玄学色彩。常见的情况有三种:

第一种是训练Loss一直在0.5以上降不下去。排查方向:数据标注是否出错、类别分布是否均衡、学习率是否过高导致震荡。如果你用的是ImageFolder加载数据,可以抽取一个batch打印标签和图片一起看,确认类别编号没有错乱。

第二种是验证集准确率非常高(比如0.98)但实际单张测试很烂。这几乎是分类项目最典型的数据泄露:train和val划分时有重复图片,或者训练时随机增强的对象被一并放进了验证集。解决办法是重建数据集划分,确保同一张图不会同时出现在两个集合中。

第三种是Loss已经收敛但准确率仍然不高。说明类别间的区分度不够,先看混淆矩阵(confusion matrix),常见的坑是两个类别互相被认错,比如“塑料瓶”和“玻璃瓶”形状颜色接近。这时候思考方向是数据而非模型:补样本、增加拍摄角度多样性,或者换成更大的模型(resnet50换成efficientnet-b4)。

我习惯的做法是训练期间每5个epoch保存一次checkpoint,文件名带上epoch号和val_acc。这样即使训练中断或过拟合,也能回溯到最佳状态;不要只依赖最后一轮的权重文件,那往往不是验证集上表现最好的。

提示:训练命令务必指定随机种子,--seed 42这种固定seed能让实验可复现。不固定种子的话,同样代码跑两次结果可能差异大,后面排查问题根本没法对照。

5. 避坑:垃圾识别分类源码最常见的6个翻车现场

5.1 解压后报ModuleNotFoundError,路径中文背锅

现象:从zip解压后,cd到项目目录执行python predict.py,系统报错提示某个模块找不到,或者路径管理混乱导致相对导入失败。

原因:zip包解压后生成的顶层目录名常带中文或特殊字符,Windows下路径解析出问题;更常见的是项目内嵌多层目录,predict.py在根目录,而models目录在另一个层级,python运行时找不到模块。

解决:先把整个项目目录移到纯英文路径下,比如D:\projects\garbage_sys,然后进入predict.py所在目录执行命令。如果还报错,用import sys; sys.path.append(os.getcwd())这种方式把根目录加进模块搜索路径。检查__init__.py是否存在于需要导包的目录中,缺失会导致相对导入失败。

5.2 权重文件损坏:预测结果全是噪声

现象:明明代码和前向流程都对,模型也能跑,但输出的概率分布几乎接近随机,分类准确率惨不忍睹。

原因:zip包在传输或网盘存储过程中文件损坏,权重文件体积不对或哈希校验没通过。PyTorch加载时有时会静默成功,但数值全是垃圾。

解决:拿到源码包后,第一时间比对压缩包和权重文件的MD5哈希值,用md5sum weights/model.pth查看哈希,和发布方给的比对(通常写在README里)。没有参考哈希时,看文件大小是否和训练记录一致,resnet50在ImageNet预训练权重约98MB,如果下载下来只有几十MB,基本是坏的。重新下载并校验后再用。

5.3 类别标签错位:验证集准确率虚高,实际推理全错

现象:训练日志显示准确率非常高,但用训练集之外的真实照片预测时,把“果皮”识别成“玻璃瓶”之类的离谱结果,而且错误方式不一。

原因:训练数据目录的顺序和模型输出的类别索引不对应。ImageFolder按字母序排列子目录,假设你的目录名是kitchen, hazardous, other, recyclable,模型输出的索引0对应哪个类别,由目录排序决定,不是由你写的类别顺序决定。如果garbage_classes.json里的映射顺序和目录排序不一致,标签就错位了。

解决:训练完成后,不要用准确率数字逆向反推类别对应关系,直接在推理脚本里打印每个索引对应的中文标签,和目录结构逐一核对。更稳妥的做法是训练前就把类别列表按目录结构的排序方式输出,确认无误后再开训:

# 验证数据集加载的类别顺序 python -c " from torchvision.datasets import ImageFolder d = ImageFolder('data/garbage_classify/train') print(d.class_to_idx) "

这段命令会打印类似{'hazardous': 0, 'kitchen': 1, 'other': 2, 'recyclable': 3}的映射,写标签文件时照这个顺序抄就不会错位。

5.4 图片尺寸和通道不一致导致的预测崩溃

现象:训练集图片一切正常,但推理时某些照片要么报错,要么结果极差;更有意思的是,OpenCV读图和PIL读图结果不一样。

原因:垃圾分类的数据来源混杂,有的图是RGB,有的是灰度图,还有的是RGBA四通道PNG。预处理流程用PIL的Image.open读入,灰度图和RGBA图不会自动转成RGB三通道,模型前向就报维度错误或数值错乱。Resize尺寸若不是模型的预期大小,模型也能跑(PyTorch的卷积对输入尺寸不完全受限),但后续全连接层可能报维度不匹配。标准做法是统一转换。

解决:在图像读取环节显式convert('RGB'),并检查模型的输入尺寸和预处理Resize是否一致:

img = Image.open(img_path).convert('RGB') # 统一转三通道 if img.mode != 'RGB': img = img.convert('RGB')

若源码里的transform用的是Resize((256, 256)),而模型定义里输入是224,中间缺了CenterCrop;或Resize用的等比缩放而非固定尺寸,都会导致训练和推理不一致。把两个参数整体对齐。

5.5 GPU显存不足:batch size调不下来时的替代做法

现象:训练脚本启动即报CUDA out of memory,把batch size从32降到16仍报错,降到4能跑但训练速度极慢,Loss波动也大。

原因:机器显存有限,加上源码默认加载全精度float32参数,模型和中间激活值占满显存。一个小容量显存显卡,载入resnet50 + 256x256输入 + batch size=32就是会超限,不是代码bug。

解决:除了调小batch size,还可以用混合精度训练。源码如果有--amp参数,直接启用;没有的话,在训练脚本里加一条自动混合精度配置:

# train.py中开启AMP自动混合精度 from torch.cuda.amp import autocast, GradScaler scaler = GradScaler() for imgs, labels in train_loader: imgs, labels = imgs.to(device), labels.to(device) with autocast(): outputs = model(imgs) loss = criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()

AMP会把部分运算降到float16,显存几乎减半,在支持半精度的显卡上效果明显。此外,把输入尺寸从256降到224(如果你的模型允许),也能省掉约25%的激活值占用。这两件事做完仍不足的,再用batch size=8配合梯度累积等效大batch效果。

5.6 实时推理耗时严重:摄像头画面帧率极低

现象:摄像头画面卡成幻灯片,帧率不到5FPS,实际运行时总是跟不上拍摄速度。

原因:源码默认图参加载到CPU推理,没有调用GPU;或者模型太大(ResNet152、EfficientNet-b5这类),在普通CPU上单帧推理耗时以秒计。

解决:先确认推理设备,torch.cuda.is_available()为真且显存足够时,代码中的map_location和.to(device)都应指向GPU。模型层面,换轻量模型是立竿见影的:MobileNetV3和EfficientNet-lite在CPU上的单帧推理通常在100ms以内,准确率损失控制在可接受范围。输出显示上,不要每一帧都做完整的前处理和模型推理,用第3章的跳帧方案把推理频率降到3~5Hz即可满足多数场景。

6. 从源码到可用系统:模型导出与部署到小机箱

6.1 把PyTorch模型导出ONNX并验证输出一致性

源码跑通、模型重训完成,接下来是要去现场部署。不能直接把.pth文件扔给现场设备,因为PyTorch在无GPU环境的依赖和版本要求太多。通用做法是导出为ONNX格式,然后在部署端用ONNX Runtime做推理。

# export_onnx.py import torch import torch.onnx # 恢复模型并加载训练好的权重 model = torchvision.models.resnet50(pretrained=False) model.fc = torch.nn.Linear(2048, 4) model.load_state_dict(torch.load('checkpoints/best_model.pth', map_location='cpu')) model.eval() # 申明ONNX导出,动态batch使运行时可以自由调整批量大小 dummy_input = torch.randn(1, 3, 224, 224) torch.onnx.export( model, dummy_input, 'garbage_model.onnx', input_names=['input'], output_names=['output'], dynamic_axes={'input': {0: 'batch'}, 'output': {0: 'batch'}}, opset_version=11 )

ONNX导出后,做一步输出一致性验证:用同一张图片分别通过PyTorch和ONNX Runtime推理,比对两者的预测结果和softmax概率分布,差异在1e-4以内就说明导出成功;如果差异大,通常是代码里有部分算子不支持ONNX导出,需要换实现或用opset更高版本。验证代码不用另写,直接写一个简短的python脚本加载模型和onnx文件各跑一次。

6.2 在边缘盒子上做推理加速的普遍做法

部署到现场设备(树莓派、Jetson Nano、工控机这种)后,推理加速有三板斧:第一,量化。把float32权重转成int8量化模型,体积缩小4倍,CPU推理速度通常能提升2到3倍,代价是准确率下降1%到3%左右。垃圾分类这种粗粒度分类场景,这个损失完全可接受。

第二,换推理框架。同样跑ONNX模型,ONNX Runtime比PyTorch在CPU上快,而OpenVINO(针对Intel CPU)和TensorRT(针对NVIDIA GPU)能在架构层面做算子融合,速度又能上一个台阶。现场是什么硬件就换什么推理后端。

第三,限制输入范围。摄像头固定在垃圾桶上方时,画面里真正需要识别的区域只有桶口附近,把ROI区域固定住,只对裁剪后的区域做推理,比在整张1920x1080画面上缩放要省得多。

这半年跟垃圾分类项目打交道下来,我的习惯是把验证脚本和推理脚本彻底分离——训练能跑、部署能跑,这是两层验收,中间隔着一个完整的环境差异。任何一个中间环节的产出(权重、类别映射、预处理参数)都不信任口头传递,全部用脚本固定下来。这样才能保证今天调通的东西,三个月后换个机器还能稳定复现。希望这些被坑过踩过的细节,帮你在自己的垃圾识别分类系统上少走一段弯路。

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

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

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

立即咨询