ARTICLE DETAIL

资讯详情

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

YOLOv8s剪枝源码深度拆解:稀疏训练、通道裁剪与微调恢复

YOLOv8s剪枝源码深度拆解:稀疏训练、通道裁剪与微调恢复 给yolov8s做剪枝这件事我一开始的想法很天真GitHub上找到一个剪枝源码clone下来把权重丢进去等它吐出一个更小的模型就行。真跑起来才发现所谓剪枝源码从来不是一个脚本跑完的事它是一整套围绕模型结构展开的工程逻辑。我花了两周从稀疏训练、通道裁剪、到微调恢复把流程完整跑通中间摔过的跟头集中在这几个地方剪完forward直接报shape不匹配、mAP掉了近十个点拉不回来、稀疏项一加loss曲线就开始抖。这篇文章不打算贴一个天上掉下来的完整项目而是把yolov8s剪枝源码中最关键的几段拆开讲清楚适合手里已经有一个能跑的YOLOv8s项目、想把它压到边缘设备上跑得更快的人。1. 动手之前先摸清YOLOv8s的冗余到底在哪里1.1 参数和算力不是平均分配的s也不是真的小YOLOv8s的参数量大致在11.2M左右640×640输入下FLOPs约28.4G。作为对比YOLOv8n只有约3.2M参数m是25.9M左右。很多人的误解是既然叫s应该已经足够小了但在真实部署场景里11.2M模型加28GFLOPs仍然是一道很高的门槛。一个原始检测工程跑在Jetson Orin Nano这类设备上TensorRT fp16下YOLOv8s大概也就勉强实时如果输入分辨率再提到1280立刻掉到十几帧。这时候要么换n系列牺牲精度要么就想办法在不伤精度的情况下压缩s剪枝因此成为很自然的选项。真正摸过模型源码之后会发现YOLOv8s的冗余并不是均匀分布的。Backbone里的C2f结构是算力大头每个C2f由1×1卷积和若干个Bottleneck组成Bottleneck里两个3×3卷积是真正的计算主体Neck部分的PAN-FPN若干个C2f同样吃算力Detect检测头虽然整体参数占比不高但每个尺度输出64nc的回归分支和nc分类分支也要占一小块。同一层卷积在不同数据集上重要性差异很大这就需要靠后续的稀疏训练去量化而不是拍脑袋决定。1.2 为什么最终选了结构化通道剪枝而不是稀疏矩阵剪枝分成两大类非结构化剪枝和结构化剪枝。非结构化剪枝把每个权重矩阵里绝对值小的元素直接置零得到的是一个松散稀疏的矩阵torch.nn.utils.prune就能做但问题在于普通推理引擎和TensorRT并不会对这种随机稀疏矩阵加速除非你有专门的稀疏算子库和硬件支持否则纯粹是纸面轰轰烈烈、实测纹丝不动。channel pruning通道剪枝是另一种思路把整张特征图里不重要的通道整个删掉。这样卷积核的维度实实在在变小了下一层的输入通道数也跟着变小模型的参数量、FLOPs、内存占用都实打实地降下来不需要特殊硬件就能在CPU、GPU上获得收益。这也是社区里主流yolov8s剪枝方案都在做的事情。通道剪枝还有一个便利条件YOLOv8的卷积基本都是Conv-BN-SiLU结构BN层的gamma系数天然可以拿来衡量通道重要性剪枝时可以直接复用训练过程中BatchNorm统计出的尺度因子不用额外设计复杂的判定网络。2. 剪枝源码的地图Ultralytics里的模块依赖比想象的复杂2.1 三个必须看明白的文件conv.py、block.py、head.py如果你和我一样是基于Ultralytics的官方YOLOv8源码来改剪枝脚本最终要面对的就是这三个文件ultralytics/nn/modules/conv.py、block.py、head.py。conv.py里定义了Conv、DWConv、Concat、autopad这一批基础组件block.py里是Bottleneck、C2f、SPPF这些积木head.py里的Detect是最后的输出头。剪枝脚本要改的每一个层都源自这三个文件。先看conv.py里的Conv模块。它的结构是Conv2d - BatchNorm2d - SiLU和经典CNN里最常见的组合一致。这个模块本身就是可剪枝单元要剪它就是根据它后面的BN gamma决定要不要保留某个输出通道然后把Conv2d的weight和BN的四个统计量全部按保留下来的序号截一遍。DWConv略微麻烦一点因为它是分组卷积每组只有1个通道剪掉一个通道等于同时改输入输出处理时要在代码里单独加判断。block.py里C2f的结构要复杂不少。C2f内部先是两个1×1卷积cv1和cv2中间夹着n个Bottleneck最后是一个Concat把多个分支拼起来。如果我们把某个Bottleneck的输出通道剪了那Concat之后的通道总数就变了cv2的输入也必须跟着改。偏偏Bottleneck里还可能有shortcutshortcut存在的时候残差相加两边的通道数必须完全一致否则直接报错。这一层套一层的依赖关系是剪枝脚本里最容易写错的地方。head.py里的Detect这里只说一句它是整个剪枝红线中最需要慎重的部分具体原因后面第四节单独讲。2.2 跨层连接让逐层剪枝不再成立我刚上手的时候写过一个很蠢的版本遍历模型所有BatchNorm2d按gamma绝对值排序把后50%通道的序号记下来然后直接对每个Conv做剪裁。跑起来就发现单独剪某一层是没问题的但只要这个层后面接了Concat或者shortcut整个通道编号就对不齐了。比如YOLOv8s的Neck里有很多Concat把上采样分支和Backbone分支拼起来两个分支必须保持相同的通道数剪枝否则Concat后的张量在通道维度上出现一半新一半旧的混乱。严格来说一个剪枝源码要处理的不是某几个层而是一张依赖图。常用的做法是给每个可剪层建一个group同一组内的层共享同一份裁剪mask。比如某次剪枝决定把第i组的输出通道剪掉30%那这一组里所有参与Concat或shortcut的层都要用同一个keep_idx来裁。社区里yolov5剪枝方案大多有类似逻辑YOLOv8因为C2f的嵌套结构更明显这个依赖关系更复杂。你没有在源码里把依赖图理顺之前所谓的剪枝脚本只能算是一个通道随机删除器。提示我第一次跑通剪枝用的不是自己硬写的依赖分析而是先在代码里打印每个C2f内部的shortcut和Concat引用手动画一遍数据流向再在剪枝函数里把同组层标记成同一个mask。过程笨但对理解源码结构很有帮助。3. 稀疏训练那一段源码让BN层的gamma替我们做标记3.1 给loss加一点点L1正则主训练循环只需要动几行剪枝不能直接对着一个普通训练好的模型做。虽然YOLOv8s本来就有一定冗余但直接按gamma排序硬剪伤害通常很大。标准做法是先做稀疏化训练在训练loss里加入对BatchNorm层gamma的L1正则让一部分通道的gamma值在训练过程中被压到接近0这些通道对应的就是模型自己认为不重要的通道。核心代码其实就一小段。不管你用的是Ultralytics自己的训练入口还是自己写的训练循环在(loss).backward()之前加上下面这段逻辑sparse_rate 0.001 # 常见范围 0.0005 ~ 0.002太大会直接压崩mAP device next(model.parameters()).device sparsity_loss torch.tensor(0.0, devicedevice) for name, module in model.named_modules(): if isinstance(module, torch.nn.BatchNorm2d): sparsity_loss torch.abs(module.weight).sum() total_loss loss sparse_rate * sparsity_loss total_loss.backward() optimizer.step()原理很简单BN层本来有weight和bias前向推理时对输入做归一化后再乘gamma加beta。对gamma的绝对值求和再乘一个小系数加进loss梯度下降时就会额外惩罚那些不必要放大或缩小的通道让它们的gamma慢慢偏向0。gamma越小的通道说明它输出的特征对后续损失影响越小剪掉它对最终预测的扰动也最小。这段里我吃过一个教训稀疏率不是越大越好。我第一次下意识设成0.01训练了150轮后loss曲线倒是很好看但验证mAP掉了将近5个百分点。后来把系数降到0.001又在warmup结束后才正式让正则项生效模型精度才基本稳住。如果你发现加了稀疏项之后训练mAP就开始明显下滑优先调小稀疏率而不是去调大。3.2 怎么看稀疏化效果gamma分布直方图比loss曲线更可靠我判断稀疏化是否到位基本不看loss曲线因为加过正则之后的loss已经不够纯净了。我习惯直接看BN层gamma的分布直方图。import numpy as np import torch import matplotlib.pyplot as plt gammas [] for name, module in model.named_modules(): if isinstance(module, torch.nn.BatchNorm2d): gammas.append(module.weight.detach().cpu().numpy().reshape(-1)) gammas np.concatenate(gammas) plt.hist(gammas, bins100) plt.yscale(log) plt.xlabel(BN gamma) plt.ylabel(通道数量(log)) plt.show()训练完稀疏模型后直方图应该出现一个明显的现象在0附近有一根很高的柱同时在远处还有一堆接近正态分布的尾巴。0附近那一簇高柱就是被正则项压到几乎为0的通道它们数量越多说明模型结构上的冗余越明显后续剪枝越从容。如果直方图看起来很居中没有明显的近0堆积说明稀疏训练还没有真正生效这时候去剪枝多半会大掉点。4. 真正动刀的核心源码通道怎么剪、层与层之间怎么联动4.1 每层剪多少全局阈值与保留比例的平衡稀疏训练结束之后拿到手的是一堆gamma值。接下来要回答剪哪些通道。社区方案里有两种主流策略全局阈值法和按层比例法。全局阈值法把所有BN层的gamma扔到一个池子里定一个阈值小于阈值的通道全部剪掉实现简单但有个致命缺陷——不同层的gamma尺度不完全一致某些层整体数值偏高全局阈值可能把它们剪成空壳。我实际采用的是按层比例保留下限的思路对每个BN层按gamma绝对值从大到小排序保留前keep_ratio * num_channels个通道同时设置一个每层最小保留比例比如无论如何每层至少保留50%通道。这能在压缩率和结构完整性之间找到平衡最后再用一个全局FLOPs目标反复调整保留比例。关键代码大概是def compute_keep_indices(bn, keep_ratio, min_ratio0.5): gamma bn.weight.data.abs() num_keep max(int(gamma.numel() * keep_ratio), int(gamma.numel() * min_ratio)) num_keep min(num_keep, gamma.numel()) _, indices torch.topk(gamma, num_keep) return indiceskeep_ratio取0.65到0.75是个比较常见的起点意味着你希望剪掉25%到35%的通道。不要上来就剪50%除非你后续做了非常强的蒸馏。剪枝率越高微调恢复难度是近似指数上升的。4.2 核心裁剪函数剪一个Conv后面所有关联层的shape都要跟着改通道剪枝代码里最核心的函数就是按keep_idx裁剪某个卷积与它配套的BN。下面是简化版逻辑对实际项目有直接参考价值def prune_conv_bn(conv, bn, keep_idx): # 裁剪输出通道 conv.weight.data conv.weight.data[keep_idx] # [keep, in, kh, kw] if conv.bias is not None: conv.bias.data conv.bias.data[keep_idx] conv.out_channels keep_idx.size(0) # 同步裁剪BN的4个统计量 bn.weight.data bn.weight.data[keep_idx] bn.bias.data bn.bias.data[keep_idx] bn.running_mean.data bn.running_mean.data[keep_idx] bn.running_var.data bn.running_var.data[keep_idx] bn.num_features keep_idx.size(0)光剪输出还不够紧随其后的那个卷积层的输入通道也要被剪。YOLOv8里多数是Conv-BN-SiLU首尾相接所以下一层通常就是一个Conv2ddef prune_conv_input(next_conv, keep_idx): next_conv.weight.data next_conv.weight.data[:, keep_idx, :, :] next_conv.in_channels keep_idx.size(0)这里有个容易忽略的坑如果下一层是nn.BatchNorm2d而不是Conv要改的是bn.num_features如果中间隔了一个Concat那不能只改Concat后面第一个Conv还得保证Concat的所有输入分支都按同一套索引被剪过。但凡有一个分支没对齐运行时报的错就是你根本想不到的尺寸不匹配。4.3 输出头、shortcut、Concat剪枝脚本里红名单与黑名单不管开源剪枝源码长什么样你都能在里面找到一份禁止剪枝或单独处理的名单通常包含三类Detect回归分支的最后一层、shortcut残差连接的必经层、Concat节点本身。Detect输出层的通道数是由DFL和类别数决定的YOLOv8每个尺度的bbox回归分支输出reg_max * 4 64个通道分类分支输出nc个通道。如果你把最后的64改成40DFL解码逻辑立刻乱套坐标预测直接废掉。分类分支同理nc由数据集决定也动不得。所以在我的剪枝脚本里Detect中的中间层可以小比例剪但输出层必须完整保留。Concat节点没有参数不需要剪但它两边参与拼接的层必须保证裁剪后维度一致我在代码里是把它们统一登记到同一个mask group里的。shortcut也一样Bottleneck里如果启用了shortcut那么这个残差分支的两端必须使用完全相同的keep_idx否则加起来数值是错位的。为什么市面上剪枝脚本越来越复杂就是因为这些隐式约束散落在模型结构的各个角落。5. 剪枝之后最常见的三场事故结构对不上、掉点、NaN5.1 剪完先别急着训用一条dummy tensor自检结构剪枝动作做完第一件事不是开训练而是让模型跑一个前向把每层实际shape打出来和剪枝前做对比。我用过一个很小的自检脚本def collect_shapes(model, h640, w640): model.eval() x torch.randn(1, 3, h, w) shapes {} hooks [] for name, module in model.named_modules(): hooks.append(module.register_forward_hook( lambda m, inp, out, nname: shapes.update( {n: tuple(out.shape) if hasattr(out, shape) else None}) )) with torch.no_grad(): model(x) for hook in hooks: hook.remove() return shapes把剪枝前后的shapes打印出来逐行diff重点看每一处Concat、Bottleneck、Detect输出。如果通道数对不上八成就是某个group的mask没有传播到位。更快的验证方式是直接把剪枝后的模型导出ONNX用onnxruntime跑一次推理ONNX导出会把一层层的常量shape检查得很严格很多Python运行时不会暴露的隐患在导出阶段就会报出来。5.2 微调不是重新训练学习率、正则、蒸馏的选择结构自检通过后才是微调。剪枝完成时模型权重已经保留了大部分有效信息不需要从头开始训练但也不能按原训练配置跑。我的经验是初始学习率降到原方案的十分之一到五分之一训练那么几十到一百轮即可。稀疏训练时加的那个L1正则到微调阶段必须关掉否则好不容易保留的通道又开始被压。微调阶段最重要的判断是掉点来源。剪枝后mAP掉2到3个点是很正常的微调后通常能回一部分。如果掉了七八个点且怎么调都回不来我会优先怀疑是不是剪到了关键结构或稀疏训练阶段gamma的分布就没有真正压好。另一个有效做法是知识蒸馏剪枝前的模型当teacher剪枝后的模型当student在训练loss里加一项蒸馏loss让student的预测分布去对齐teacher。这比单纯re-training回血快得多尤其在目标检测任务上teacher给出的软标签能稳得住那些被剪掉通道对应的语义信息。下面是我常用的微调参数参考项目推荐设置说明初始学习率原方案lr的1/10到1/5太大容易跳过细粒度结构微调轮数30到100轮取决于数据量和掉点幅度稀疏正则关闭不能继续压gamma蒸馏可选掉点超过3个点建议加5.3 mAP掉点后的排查顺序与NaN的常见来源如果微调效果还是不好三个方向值得逐一排查。第一回看稀疏训练的直方图0附近没有明显的堆积说明模型根本没在稀疏化训练里学到哪些通道不重要这时需要调大正则系数或延长稀疏训练轮数而不是在微调阶段死磕。第二检查隐藏危险层有没有被误剪比如某个shortcut分支的mask不一致、Concat分支的通道编号错位。第三BN统计量需要重新校正剪枝后模型结构突变BN的running_mean和running_var可能已经失真我会在微调训练前用一小批训练数据单独跑几轮trainingFalse的前向来刷新统计量这能直接改善推理阶段的初始表现。至于NaN多数发生在剪枝后直接开训且学习率偏大的情况下BN在某些通道被剪掉后方差统计出现极端值导致梯度在回传时爆炸。处理办法其实很朴素先把学习率压到更低加上梯度裁剪同时用一小批数据校正BN统计量再正式训练。如果NaN还出现就去检查是不是某个被保留的BN通道确实存在数值异常比如剪完后的running_var出现负数或者inf。这几个问题本质上是结构正确性和统计量正确性两个层面排查时要分开处理不要混在一起调参数。6. 剪完怎么验证收益从FLOPs表到端到端延迟6.1 别只看FLOPs把参数量、显存、延迟放在一张表里看很多人汇报剪枝成果习惯只报FLOPs减少百分之多少但在真实工程里这个数字参考意义有限。我测完一个剪枝模型会固定好输入分辨率、批大小、推理后端把参数量、模型文件大小、FLOPs、显存占用、端到端延迟打在同一张表里看。下面是我某次剪枝前后在TensorRT fp16 RTX 3060上测的示例数据不同环境差距很大仅供参考指标剪枝前剪枝后参数量11.2M6.5M模型文件大小22.4MB13.1MB640输入FLOPs28.4G19.3GTensorRT fp16延迟约6.8ms约5.1ms从表里能看出一个关键现象FLOPs降了约32%但端到端延迟只降了约25%。原因在于端到端推理还有NMS、数据预处理、kernel launch和内存带宽这些固定开销它们不受剪枝率线性影响。如果目标场景是CPU或嵌入式端通道剪枝的收益通常比GPU上更明显因为CPU上的计算密集度瓶颈更突出。6.2 剪枝、蒸馏、量化部署前的组合拳怎么打通道剪枝很少单独上场。我在实际项目里经常用的组合是先做蒸馏训练获得一个高精度teacher再对student模型做稀疏训练和通道剪枝剪完推出ONNX然后接TensorRT的fp16或int8量化。三者顺序上剪枝和量化可以先后进行但要注意精度叠加损失剪枝掉3个点int8量化再掉1到2个点最后的模型可能已经到边缘可接受范围。如果你对精度极其敏感可以先只做fp16量化留出剪枝的精度预算。还有一个个人习惯剪枝前先确认目标平台的上限。如果真正约束你的是显存而不是计算时间那么压缩模型文件大小和减少channel数是主要诉求如果约束是延迟则要考虑剪枝在目标推理引擎里带来的实际加速曲线。我见过有人费大力气把模型从11M压到5M最后在GPU上因为小算子调度碎片化严重延迟反而更差的情况。剪枝源码能帮你把FLOPs降下来但能不能落地成真收益一定取决于你最终跑在哪个硬件、哪个推理引擎、哪个精度模式。最后分享一个我自己保留至今的习惯每次剪枝前先把原始模型的mAP和正常训练的loss曲线存个档剪枝后微调时每次都拿回来对照。如果微调到后期mAP差距始终在可接受范围内说明这刀砍得合理如果差距一直追不回来我会回到稀疏训练阶段把稀疏系数调低重跑而不是在微调阶段硬扛。剪枝这把刀砍下去容易真正难的是知道砍哪里、砍多深。希望这篇文章里从源码角度拆出来的这些细节能让你少踩几个我踩过的坑。
返回列表