news 2026/9/14 5:55:24

基于StemBlock与ShuffleNet的YOLOv5轻量化垃圾检测改进

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
基于StemBlock与ShuffleNet的YOLOv5轻量化垃圾检测改进

简介:面向高校人工智能、电子信息、自动化等专业学生及毕业设计、课程设计和竞赛项目研发人群,本资源是一套可实际运行的垃圾分类检测系统。项目基于YOLOv5进行改进,引入Stemblock与Shufflenet结构,在轻量化部署与检测精度之间做了针对性平衡,覆盖数据集、模型配置、训练流程和项目说明文档,支持直接用于课设、毕设或项目初期演示。资源包共1117个文件,约71.09MB,以txt标注文件、jpg/jpeg图像样本为核心,配合py训练脚本、yaml模型配置、pt权重文件和ipynb示例,并附garbage_classification.db等辅助资料,目录结构清晰,可快速定位数据预处理、模型训练与检测推理等模块。已有54人学习下载,适合希望深入理解YOLOv5改进思路并快速搭建实际项目的初学者,也便于在此基础上二次开发,扩展其他检测场景。

1. 垃圾分类检测为什么要动YOLOv5的Backbone

垃圾分类检测系统这几年几乎成了深度学习课设里的“压轴题”,但要把它做到能演示、能答辩,并不只是跑通一个YOLOv5那么简单。垃圾样品类别多、外形不规则,一次性餐盒、易拉罐、塑料袋在画面里经常小且互相遮挡,直接拿默认的YOLOv5s去训练,mAP往往卡在一个不上不下的位置。这个项目做的改进很直接:把YOLOv5的主干CSPDarknet换成StemBlock加ShuffleNet的组合,参数量更小、推理更快,对折叠、半遮挡的小目标反而更稳。它适合计算机视觉方向的学生做课程设计或毕业设计,也适合想评估轻量化Backbone对检测精度影响的一线开发。下面按“为什么这样改、在哪改、怎么训练、怎么排查、怎么部署”把完整通路过一遍。

2. 原版CSPDarknet的瓶颈与Stem、ShuffleNet的设计逻辑

2.1 原版Backbone在垃圾分类场景下的三个问题

YOLOv5原版主干从6x6卷积起步,接若干C3模块和SPPF,设计目标是兼顾ImageNet分类精度和推理速度。但在垃圾分类这类细粒度小目标任务里,它有三个明显短板。

第一,前两次下采样太快。垃圾图像里一次性纸杯、塑料瓶盖这类目标本身只占几十个像素,原版主干在前两层就做了4倍下采样,浅层的细节纹理大量丢失,后面网络只能靠语义信息“猜”边界。第二,C3模块的瓶颈结构依赖大量3x3卷积,中间通道数翻倍又压缩,对边缘设备不友好。第三,原版Focus切片虽然把空间信息拆到通道里,但在实际部署时对TensorRT和ONNX的算子优化并不友好,有些设备上反而慢。

所以在不改检测头的前提下,把Backbone换成StemBlock加ShuffleNetV2成为这个项目的首选方案。StemBlock负责更温和的下采样,ShuffleNet负责用分组卷积压低计算量,两者互补,改动也只集中在models目录下。

2.2 StemBlock:两条下采样支路并行的轻量入口

StemBlock最初出现在CSSNet里,核心思路是用一个3x3 stride=2卷积和一条“maxpool + 1x1降维 + 3x3提特征”的支路并行做下采样,最后concat再融合。相比单一卷积下采样,它能同时保留两种感受野下的信息。

# models/common.py 中新增的 StemBlock class StemBlock(nn.Module): def __init__(self, c1, c2, k=3, s=2): super().__init__() # 3x3 stride=2 卷积支路,通道从 c1 升到 c2 self.conv1 = Conv(c1, c2, k=3, s=2) # 2x2 stride=2 最大池化支路,不增加参数 self.maxpool = nn.MaxPool2d(kernel_size=2, stride=2) # 池化支路后续先降通道,再还原 self.conv2 = Conv(c2, c2 // 2, k=1, s=1) self.conv3 = Conv(c2 // 2, c2, k=3, s=1) # 融合后把通道压回 c2,避免直接翻倍 self.conv4 = Conv(c2 * 2, c2, k=1, s=1) def forward(self, x): x1 = self.conv1(x) x2 = self.maxpool(x) x2 = self.conv2(x2) x2 = self.conv3(x2) out = torch.cat([x1, x2], dim=1) return self.conv4(out)

这段代码里,conv1和maxpool的输出宽高一致,分别是原图1/2,所以concat没有对齐问题。conv4的1x1卷积把两路拼出来的2*c2通道压缩回c2,控制后续ShuffleNet的输入规模。注意这里复用的是YOLOv5自带的Conv类,它内置BN和SiLU激活,不再额外加激活层。

2.3 ShuffleNetV2单元:分组卷积、通道重排与内存访问成本

ShuffleNetV2的设计依据不是FLOPs而是直接测内存访问成本MAC。四个原则:输入输出同通道时MAC最小;分组数过大会增加访存;碎片化结构对并行不友好;逐元素运算也要算时间。落到具体模块上,就是stride=1时一半通道走恒等映射,另一半走卷积,最后用channel shuffle交换两组信息。

def channel_shuffle(x, groups): b, c, h, w = x.shape x = x.view(b, groups, c // groups, h, w) x = x.transpose(1, 2).contiguous() return x.view(b, -1, h, w) class ShuffleV2Block(nn.Module): def __init__(self, inp, oup, stride): super().__init__() self.stride = stride mid = oup // 2 if stride == 1: # 左支路恒等,保证输入输出通道相同 self.branch1 = nn.Identity() else: # 左支路用 depthwise 卷积做 2 倍下采样 self.branch1 = nn.Sequential( Conv(inp, inp, k=3, s=stride, g=inp), Conv(inp, mid, k=1, s=1), ) self.branch2 = nn.Sequential( Conv(inp, mid, k=1, s=1), Conv(mid, mid, k=3, s=stride, g=mid), Conv(mid, mid, k=1, s=1), ) def forward(self, x): if self.stride == 1: x1, x2 = torch.chunk(x, 2, dim=1) else: x1, x2 = x, x out = torch.cat([self.branch1(x1), self.branch2(x2)], dim=1) return channel_shuffle(out, 2)

关键参数是mid = oup // 2。stride=1时输入输出通道一样,stride=2时两条支路各输出mid,concat后正好翻倍。这里用YOLOv5的Conv类替代论文里的原始卷积,等于在depthwise卷积后面也加了BN和SiLU,实测在YOLO这种带大检测头的结构里比原版更稳。

3. 改造YOLOv5网络:从common.py到yolo.py的完整接入

3.1 项目源码里和改动相关的文件

拿到源码包后,先别急着跑train.py,把目录结构看清楚。这个项目里和网络改动强相关的文件基本都在models和data目录下。

文件/目录作用
models/common.py网络基础组件,StemBlock和ShuffleV2Block加在这里
models/yolo.pyparse_model解析yaml并组装模型,必须同步注册新模块
models/yolov5_stem_shuffle.yaml改进后的Backbone结构配置
data/garbage.yaml垃圾分类数据配置,指定图片路径和类别数
run_detect.bat一键运行检测的批处理文件
garbage_classification.dbSQLite数据库,存放类别信息与识别记录

另外,项目数据目录下能看到train2017.cache,这是YOLOv5第一次跑训练时生成的图片索引缓存。如果你修改了图片路径或增删了图片,建议删掉这个缓存再训练,否则会一直读到旧的索引。

3.2 在yolo.py的parse_model里注册新模块

models/yolo.py的parse_model函数是整个网络装配的中枢。它逐行读yaml里的backbone和head,遇到没见过的模块名会直接抛错。所以要加一行分支,把StemBlock和ShuffleV2Block的构造参数解析逻辑补充进去。

# models/yolo.py 的 parse_model 里,找到模块类型判定区 if m in {Conv, GhostConv, C3, SPPF, ...}: c1, c2 = ch[f], args[0] if m is C3: args = [c1, c2, *args[1:]] ... elif m in {StemBlock, ShuffleV2Block}: # 这两个模块的构造签名是 (inp, oup, stride) # 必须把上一层的输出通道 ch[f] 作为 inp 传进去 c1, c2 = ch[f], args[0] args = [c1, c2, *args[1:]]

如果不加这个分支,parse_model默认把ch[f]当成模块的c1,但StemBlock的构造顺序是inp、oup、stride,参数对不上,初始化阶段就会报错。注册完之后,记得在yolo.py文件头的import区域把这两个新类拉进来。

3.3 用yaml拼装改进版Backbone

项目里的models/yolov5_stem_shuffle.yaml就是改完的完整结构,Backbone部分核心如下。

# models/yolov5_stem_shuffle.yaml 中 backbone 部分 backbone: [[-1, 1, StemBlock, [128]], # stride=4, 通道128 [-1, 1, ShuffleV2Block, [256, 2]], # stride=8, 给P3 [-1, 3, ShuffleV2Block, [256, 1]], [-1, 1, ShuffleV2Block, [512, 2]], # stride=16, 给P4 [-1, 7, ShuffleV2Block, [512, 1]], [-1, 1, ShuffleV2Block, [1024, 2]], # stride=32, 给P5 [-1, 3, ShuffleV2Block, [1024, 1]], [-1, 1, SPPF, [1024, 5]]]

和原版yolov5s.yaml对比,原来第一个6x6卷积和C3全被替换,SPPF保留在最后用来扩大感受野。注意每个ShuffleV2Block的第二个参数是stride,不是通道数。第一个数字才是输出通道,输入通道由parse_model自动从上一层拿。给脖子head用的P3、P4、P5分别对应第四条、第六条、第七八条输出,head部分Concat的from要改成对应的层索引,这个在项目里已经调好,自己改yaml时要格外小心。

3.4 预训练权重怎么处理

Backbone结构变了,直接用yolov5s.pt加载会打印一堆“Transferred 100/362 items”之类的警告,说明卷积层权重对不上,主干部分相当于是随机初始化。对这个项目,建议冷启动训练,也就是--weights '',让整个模型从头学。如果你的数据量不到几千张,也可以先锁住检测头前30轮,只训Backbone。

4. 垃圾分类数据集的整理与训练流程

4.1 数据格式与标注文件

项目里的图片大多是哈希命名的jpg,和它们同名的txt是YOLO格式标注。每行一个目标,格式是class x_center y_center width height,前三个值都是相对图像宽高的0到1小数。比如2 0.450 0.620 0.120 0.085,表示类别索引2的目标,中心点在图像45%、62%的位置,宽高占整图的12%和8.5%。

类别清单在哪看?项目里有个garbage_classification.db,用SQLite查一下就行。

sqlite3 garbage_classification.db ".tables" sqlite3 garbage_classification.db "SELECT * FROM classes LIMIT 10;"

一般这个库里会存一张类别表和若干识别记录表。如果你打不开db,去数据目录找classes.txt,一行的类别名顺序就是标注文件里的索引顺序。改数据集时这两个文件必须保持一致。

4.2 划分训练集与验证集

数据准备阶段最容易踩的坑是:图片和标签文件名对不上,或者某些类别只在训练集里出现。先跑一段脚本校验并划分。

# split_data.py import os import random import shutil root = 'datasets' imgs = [f for f in os.listdir(os.path.join(root, 'images')) if f.endswith('.jpg')] # 固定随机种子,保证多次划分结果一致 random.seed(3407) random.shuffle(imgs) split = int(len(imgs) * 0.9) train_files = imgs[:split] val_files = imgs[split:] for f in train_files: shutil.move(os.path.join(root, 'images', f), os.path.join(root, 'train2017', 'images', f)) shutil.move(os.path.join(root, 'labels', f.replace('.jpg', '.txt')), os.path.join(root, 'train2017', 'labels', f.replace('.jpg', '.txt')))

这段代码把90%的图片划给训练集,10%留给验证集,并且把标签同步移动。random.seed固定后每次跑结果一样,答辩时方便说明数据划分逻辑。

4.3 训练命令与超参数选择

进入项目根目录,环境配好之后,训练命令长这样:

python train.py \ --data data/garbage.yaml \ --cfg models/yolov5_stem_shuffle.yaml \ --weights '' \ --batch-size 16 \ --epochs 120 \ --img 640 \ --workers 4

data/garbage.yaml里主要改三处:path指向数据根目录,train和val分别是对应的图片目录,nc改成你自己的垃圾类别数量。数值参数建议如下,显存不够的时候按这个顺序降:先调batch-size到8,再调--img到512。

参数建议范围说明
batch-size8~32显存占用和收敛稳定性首要调这个
epochs120~200数据量大时150往上,小数据量120就够
img640或896小目标多就896,速度快就640
patience20早停轮数,防止过拟合后浪费时间
workers4~8Windows下建议设为0,避免DataLoader卡死

4.4 训完怎么读日志和选权重

训练完去runs/train/exp下看results.csv,第一列是epoch,后面依次是train loss、val loss、P、R、mAP50、mAP50-95。选权重别只看last.pt,优先跑test.py验证best.pt。这里要特别提醒:如果训练集loss一直降但val的mAP50连续20轮不涨,多半是过拟合,去data/hyps里把正则化系数调大,或者用更大的--img做数据增强。

5. 检测推理、界面与常见问题排查

5.1 run_detect.bat到底干了什么

项目里的run_detect.bat本质是一个封装好的detect.py调用。

@echo off title Garbage Detection python detect.py ^ --weights runs/train/exp/weights/best.pt ^ --source data/images ^ --conf 0.35 ^ --iou 0.45 pause

两台设备拿到同一份代码,一个能跑一个跑不了,差异通常在torch和torchvision版本。YOLOv5对torch版本比较敏感,我一般用torch 1.13到2.0之间的版本,CUDA就装11.7或12.1。环境配置这一步卡住的话,优先看requirements.txt里torch那行的等号约束。

5.2 置信度阈值和NMS参数怎么配合

detect.py推理时的两个关键参数是--conf和--iou。

参数实际含义调参方向
conf 0.35低于0.35的框直接丢弃误检多就调高,漏检多就调低
iou 0.45NMS去重时的重叠容忍度密集堆放场景调低到0.35
max_det 300全图最多保留的检测框数大批量流水线场景可以调小

垃圾分类里易拉罐和纸盒经常摞在一起,IOU阈值建议先跑一批图看看,框重叠严重就降到0.4以下,不要把默认值一把梭。

5.3 三个高频故障点

  • 训练时报AssertionError: train: No labels in xxx:说明images和labels目录没配对,回去查4.1的格式。
  • 加载权重报unexpected key成片出现:backbone结构不匹配,确认--cfg用的是yolov5_stem_shuffle.yaml,不是原版yolov5s.yaml。
  • 显存溢出:把batch-size降到4,再把--img从640降到512,同时把--workers设为0,Windows下worker线程也会吃显存。

提示:如果跑detect.py不报错但检测框全空,先用项目自带的测试图跑一次。测试图都空,检查--weights路径是否指向best.pt;测试图正常、自己的图空,大概率是数据分布差太远,考虑加背景类或者做domain adaptation。

6. 最后的技巧:把改进版YOLOv5部署成可演示的系统

6.1 用Flask包一层检测接口

答辩和课程演示时,命令行一张张出图不够直观。常见做法是写一个最简单的Flask服务,前端传图片,后端返回JSON。

# app.py import cv2 import numpy as np from flask import Flask, request, jsonify import torch model = torch.hub.load('', 'custom', path='runs/train/exp/weights/best.pt', source='local') app = Flask(__name__) @app.route('/predict', methods=['POST']) def predict(): f = request.files['image'] img = cv2.imdecode(np.frombuffer(f.read(), np.uint8), cv2.IMREAD_COLOR) results = model(img) dets = [] df = results.pandas().xyxy[0] for _, row in df.iterrows(): dets.append({ 'class': row['name'], 'conf': round(float(row['confidence']), 4), 'bbox': [round(row['xmin']), round(row['ymin']), round(row['xmax']), round(row['ymax'])] }) return jsonify({'detections': dets}) if __name__ == '__main__': app.run(host='0.0.0.0', port=8080)

torch.hub.load的path指向项目里训练好的best.pt,source='local'表示不联网。返回的bbox坐标是整数像素值,前端可以直接画框。

6.2 把识别记录写进SQLite

项目自带garbage_classification.db,训练阶段存的是类别表,推理阶段完全可以复用,把每次检测结果写进一个run_log表。用Python的sqlite3标准库就能完成,不需要额外装ORM。

import sqlite3 import time conn = sqlite3.connect('garbage_classification.db') cur = conn.cursor() for det in dets: cur.execute( 'INSERT INTO run_log (class, conf, x1, y1, x2, y2, ts) ' 'VALUES (?, ?, ?, ?, ?, ?, ?)', (det['class'], det['conf'], *det['bbox'], time.time()) ) conn.commit() conn.close()

这里把det的bbox四个值用*展开,直接匹配x1、y1、x2、y2四个占位符。表结构里ts存unix时间戳,按天聚合就能画各类垃圾的识别次数趋势图,课设答辩时是个加分项。

6.3 验证改进收益的三个硬指标

最后,用thop统计FLOPs和参数量,拿原版yolov5s.yaml同条件对比。

from thop import profile from models.yolo import Model # 统计改进后的模型 net = Model('models/yolov5_stem_shuffle.yaml', nc=3) flops, params = profile(net, inputs=(torch.randn(1, 3, 640, 640),)) print('FLOPs = %.2fG, Params = %.2fM' % (flops / 1e9, params / 1e6))

把models参数换成原版yolov5s.yaml再跑一次,对比FLOPs、参数量、单张推理耗时和验证集mAP。项目说明里给的结论是速度和精度优于原版,但你自己的数据集上要以实际输出为准。这四个数字记录到实验表格里,比任何描述都有说服力。

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/14 5:55:18

OpenCV传统车牌识别全链路实现:HSV定位+投影分割+SVM分类

简介:本资源是一套完整的基于OpenCV的Python车牌识别系统源码,面向计算机、人工智能、自动化等专业的学生与初学者,适用于毕业设计、课程大作业及计算机视觉入门实践。项目已通过答辩评审(得分98分),代码经…

作者头像 李华
网站建设 2026/9/14 5:53:03

Vue2+Element UI实现可拖拽甘特图:日期坐标换算与拖拽闭环

简介:面向Vue2开发者的可拖拽甘特图组件源码,基于Element UI实现,专门解决排期、项目管理场景中时间块拖拽调整的交互需求,避免付费插件和英文文档带来的接入成本。压缩包共21个文件,包括7个JS逻辑文件、6个Vue组件、2…

作者头像 李华
网站建设 2026/9/14 5:52:51

Rust函数编程:从基础到高级特性解析

1. Rust函数基础概念与核心特性Rust作为一门现代系统编程语言,其函数设计融合了安全性、性能与表达力三大核心优势。与C/C等传统系统语言不同,Rust函数在编译阶段就通过所有权机制消除了数据竞争和内存安全问题。一个基础的Rust函数定义如下:…

作者头像 李华
网站建设 2026/9/14 5:51:49

四模态融合课堂感知系统:情绪+表情+姿态+人脸协同分析

简介:本资源是一个基于多模态AI技术的智能教室系统实现方案,面向教育信息化开发者、计算机视觉方向学习者及智慧校园建设实践者,聚焦课堂行为分析与考试监管场景,解决学生专注度量化、动态考勤、情绪状态识别、异常姿态监测及作弊…

作者头像 李华