ARTICLE DETAIL

资讯详情

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

YOLOv8s剪枝源码实战:通道剪枝与推理加速

YOLOv8s剪枝源码实战:通道剪枝与推理加速 模型体积大、推理慢部署到边缘设备总被嫌弃这是我在跑YOLOv8s项目时最头疼的问题。后来靠剪枝解决了实测大概能砍掉30%-50%的参数推理速度提升明显精度还能维持在可接受范围。这篇就围绕yolov8s模型剪枝的源码实现展开掰开揉碎讲讲我是怎么做的从原理、工具选型到具体代码和微调过程适合正在做模型压缩、准备把检测模型部署到嵌入式设备上的朋友参考。1. 剪枝前的准备与整体思路1.1 为什么选择YOLOv8s剪枝而不是直接换更小模型YOLOv8系列本身就有n/s/m/l/x几个尺寸。如果只是追求更小体积直接换yolov8n就行了但实际体验下来有个问题小模型往往需要从头训练数据量不够或者训练时间不足时精度掉得比剪枝还狠。而剪枝是在你已经训练好的大模型基础上做文章相当于“保留知识删除冗余”它在合适的微调策略下能比直接换小模型保留更多语义信息。我手头有个项目要在Jetson Orin Nano上做实时检测原版的YOLOv8s参数约11.2M实际yolov8s参数量大概11.2M左右FLOPs也偏大部署后帧率不达标。一开始试着换成yolov8n精度掉了6个点业务方不接受。后来改成对yolov8s做结构化剪枝把通道砍掉一部分精度只掉了2个点左右帧率却提了40%以上效果立竿见影。所以一句话总结剪枝适合你已经有一个可用的模型但希望它更小更快同时不想牺牲太多精度的场景。如果你是从零开始且数据量不大老老实实训练小模型往往更省事。1.2 剪枝算法核心概念结构化与非结构化剪枝剪枝本质上就是去掉网络里不重要的权重或通道。但“怎么去掉”决定了后续能不能真正加速。非结构化剪枝把权重矩阵里绝对值较小的单个元素置零。这种做法会让模型变成稀疏矩阵需要专门的稀疏库才能提速而且硬件对不规则稀疏的支持很有限。在PyTorch环境下非结构化剪枝后模型文件确实能变小但推理速度几乎没变化。所以很多人做完非结构化剪枝觉得“白干一场”原因就在这里。结构化剪枝直接删除整个通道、滤波器或层。剪完后模型结构变了通道数变少推理时矩阵运算规模自然缩小配合GPU或CPU都能直接受益。这是实际部署中最常用的方案YOLOv8s剪枝一般指的也是结构化剪枝。结构化剪枝又分为两类基于BN层gamma系数的剪枝和基于通道重要性的剪枝。前者比较直观——BN层里的gamma值学习的是每个通道的缩放因子gamma接近0的通道意味着这个通道对输出贡献很小删掉对性能影响不大。YOLOv8s的C2f模块、卷积层后面都带着BN层天然适合这种做法。1.3 剪枝工具选型torch_pruning、NNI还是手写源码我自己在YOLOv8s上尝试过三种路线工具优点缺点torch_pruning支持结构化剪枝自动处理依赖关系接口简单对YOLOv8的C2f、SPPF等自定义模块需要手动适配NNI阿里开源功能丰富支持多种剪枝算法太重了依赖一堆组件学习成本高手写剪枝源码完全可控能针对模型结构做定制工作量大各种依赖关系容易出bug综合对比后我最终选择了torch_pruning作为主体框架但我没有直接用它的高层API而是基于它提供的底层剪枝工具手写了一部分源码。这样既能利用它处理层依赖的逻辑又能针对YOLOv8s的检测头、C2f模块做灵活调整。torch_pruning的核心优势在于它会自动分析网络层的依赖关系比如你剪掉一个卷积的某些通道它知道后面的BN层、激活层、下一个卷积层也要跟着剪掉。这个自动依赖处理对于YOLOv8s这种结构复杂的模型太重要了如果纯手写依赖处理要维护一张巨大的图很容易漏。2. 源码级拆解YOLOv8s剪枝实现要点2.1 模型结构解析与可剪枝层识别先看YOLOv8s整体结构。和v5相比yolov8s用C2f模块替代了原来的C3保留了SPPF空间金字塔池化检测头变成了解耦头Separated head也就是分类分支和回归分支分开。拿到模型后第一步是遍历model.named_modules()把可剪枝层列出来。torch_pruning里可剪枝的层通常包括Conv2d、BatchNorm2d、Linear等。但YOLOv8s里面有些层它默认不支持比如C2f里的卷积层虽然也挂着Conv2d但因为这些模块是自定义的直接剪可能破坏模块内部结构。所以我在源码实现里做了一层筛选并且对C2f内部单独处理。我自己写了一个分析函数输出模型的层结构概览def analyze_model_structure(model): conv_count 0 bn_count 0 for name, module in model.named_modules(): if isinstance(module, torch.nn.Conv2d): conv_count 1 print(f{name}: Conv2d, in_ch{module.in_channels}, out_ch{module.out_channels}) elif isinstance(module, torch.nn.BatchNorm2d): bn_count 1 print(fTotal Conv2d layers: {conv_count}, BN layers: {bn_count})这样跑一遍后你会发现YOLOv8s里的卷积层主要分布在Backbone、Neck和Head三个部分。但要注意yolov8s的解耦头里分类分支和回归分支是两个并行的Conv块直接统一剪可能会影响两个分支的一致性。我一般把Backbone和Neck作为剪枝重点对Head只轻微剪或者不剪因为检测头的精度太敏感了。另外YOLOv8s有一个细节整个模型没有全局的BatchNorm在head上其实是有的只是head部分的卷积后接BN会被torch_pruning识别。不过从实践来看剪head收益不大风险倒是很高。所以我在源码里对head部分的层加了白名单不做剪枝。2.2 基于torch_pruning构建剪枝流程torch_pruning的核心API是tp.prune_conv2d或tp.prune_module以及从模型图中获取依赖关系。我的做法是基于它提供的DepGraph依赖图来获取每个层对应的影响层组。核心流程分三步加载训练好的YOLOv8s权重。构建DepGraph调用tp.prune_conv2d对指定的conv层进行剪枝。保存剪枝后的权重并重写模型结构因为通道数变了。下面是我在源码里实现的关键片段import torch import torch_pruning as tp from ultralytics import YOLO def prune_yolov8s(model_path, ratio0.3): # 加载YOLOv8s模型这里用ultralytics的YOLO类 model YOLO(model_path) net model.model.train() # 拿到底层nn.Module net.eval() # 使用随机输入跑一遍让torch_pruning识别依赖图的输入形状 example_inputs torch.zeros((1, 3, 640, 640)) # 构建依赖图 dep_graph tp.DependencyGraph().build_dependency(net, example_inputsexample_inputs) # 收集需要剪枝的卷积层排除head和某些关键层 prunable_layers [] for name, module in net.named_modules(): if isinstance(module, torch.nn.Conv2d): # 跳过detect层相关卷积通过name关键字进行过滤 if detect in name or cv2 in name or cv3 in name: continue prunable_layers.append((name, module)) # 计算每个卷积层的通道重要性和剪枝计划 # 这里用一个简化的策略根据BN层的gamma绝对值进行排序 for idx, (name, conv) in enumerate(prunable_layers): # 找到这个conv对应的BN层通常紧跟在后面或通过dep_graph找到 bn_node dep_graph.get_related_nodes(conv)[0] bn_module net.get_parameter(bn_node.target) if isinstance(bn_node.target, str) else None # 获取这个卷积层的输出通道数 out_channels conv.out_channels # 计算要剪除的通道数 num_prune int(out_channels * ratio) if num_prune 0: continue # 找到对应的BN层获取gamma值 for sub_name, sub_module in net.named_modules(): if sub_module is bn_node.module: gamma sub_module.weight.data.abs().detach().cpu().numpy() break # 选择gamma值最小的通道进行剪枝 prune_indices torch.tensor(gamma).argsort()[:num_prune] # 执行剪枝 pruning_plan dep_graph.get_pruning_plan(conv, tp.prune_conv_out_channels, idxsprune_indices) pruning_plan.exec() # 重新构建模型结构并保存 net_fused net # 保存权重需要保存为YOLO格式需要调整model的state_dict # 这里保存为原始的pytorch权重后续再用ultralytics的模型类加载 torch.save(net.state_dict(), yolov8s_pruned.pth)注意这个代码片段是一个简化版本。实际源码里我做了很多兼容处理过滤head部分具体过滤条件要看yolov8s源码里的层命名。在ultralytics的model.py里检测头是model.model[-1]也就是Detect类它的内部有cv2和cv3这些卷积都不该被剪。对于C2f模块因为里面用到了split、cat等操作直接剪channel会导致后续拼接维度对不上。我用的办法是让torch_pruning自动处理依赖它会根据cat操作识别出需要同步剪枝的层。SPPF层里有比较大kernel的卷积也需要注意依赖关系。不过torch_pruning对MaxPool2d和Cat的处理已经很成熟了实测下来能正确传播。2.3 剪枝比例与通道选择策略剪枝比例怎么定不能拍脑袋。我一般会先做一个“剪枝敏感性分析”尝试几个不同比例比如0.2、0.3、0.4、0.5每个比例下剪完后做个短时间微调再在验证集上测精度和速度找最佳拐点。这里的关键是“最佳拐点”——通常随着剪枝比例上升模型体积和延时下降但精度在某个点后急剧下跌。我们做项目时tolerance就是业务允许的精度损失上限比如不超过2%在这个预算里尽可能剪多。通道选择策略上常见的有两种基于BN gammagamma越小越不重要。基于激活值幅度统计每个通道输出的平均绝对值越小越不重要。我的经验是BN gamma在实际操作中更稳定。因为它已经被模型训练时优化过了能反映出通道在当前流形中的贡献度。激活值统计需要跑一份数据来统计很容易受到batch选择的影响而且检测模型的特征图往往有较多spatial信息统计量波动很大。举一个实际例子我在一个交通标志检测任务里对比过两种策略策略原始mAP50剪枝后mAP5030%比例精度损失随机剪0.8220.7645.8%BN gamma剪0.8220.8061.6%激活值统计剪0.8220.7982.4%所以BN gamma是最推荐的方案。如果你的模型没有BN层那就得用激活值统计或者别的启发式方法但YOLOv8s基本都有BN直接放心用。3. 实操过程与微调方案3.1 剪枝后的模型微调配置剪完枝的模型不能直接用因为删除通道后剩余的权重是“残缺”的精度肯定会掉。这时候需要用训练集做微调fine-tune让模型重新适应新的网络结构。微调做得好精度能恢复到接近原始水平。微调和重新训练不一样学习率必须调低。我通常用的是初始学习率0.0005左右是正常训练的一个数量级以下。训练轮次选择上剪枝比例在30%以下时30个epoch就够比例超过50%可能需要80到100个epoch才能恢复。我微调时使用的配置大概是这样的# fine_tune.yaml task: detect mode: train model: yolov8s_pruned.pth # 剪枝后的结构权重 data: my_dataset.yaml epochs: 50 lr0: 0.0005 lrf: 0.05 batch: 32 imgsz: 640 optimizer: AdamW workers: 8 patience: 5这里有个关键点剪枝后的模型结构已经变了最稳妥的方式是让ultralytics从剪枝后的网络结构重新构建模型再加载权重。我在源码里提供了一种做法from ultralytics import YOLO # 直接传入剪枝后的模型文件 model YOLO(yolov8s_pruned.yaml) # 这是根据剪枝结果导出的新模型结构文件 model.load(yolov8s_pruned.pth) # 加载剪枝后的权重 model.train(datamy_dataset.yaml, epochs50)不过导出新yaml结构比较麻烦因为你得知道每一层的out_channels变了多少。我在源码中有个函数自动遍历剪枝后的net将各Conv层输出通道变化记录到字典然后根据原yaml生成新的yaml。这个思路和ultralytics本身的结构解析是兼容的。具体实现思路def generate_pruned_yaml(original_yaml, channel_changes, output_yamlyolov8s_pruned.yaml): # original_yaml是yolov8s.yaml的路径 # channel_changes是一个字典记录了每个层名字对应的新通道数 with open(original_yaml, r) as f: model_cfg yaml.safe_load(f) # 修改backbone/head的channel配置 # 这个需要根据层名字和yaml的索引对应起来比较复杂 # 我实际用的是另一个技巧直接保存nn.Module结构然后用torch.save保存整个模型对象 # 更简单的做法推荐直接保存整个模型类 torch.save({model: net, state_dict: net.state_dict()}, yolov8s_pruned_full.pth)如果你怕麻烦我的建议是剪枝完成后不要试图重建yaml直接保存整个nn.Module对象。在ultralytics中可以通过下面的方法加载剪枝后的模型import torch ckpt torch.load(yolov8s_pruned_full.pth) net ckpt[model] # 将net包成YOLO类需要一点转换或者直接用net进行推理但注意ultralytics的YOLO类并不直接支持传入一个自定义nn.Module所以如果要在train模式下微调我最终是采用把剪枝后的网络结构通过写yaml的方式恢复这是ultralytics官方支持的方式。具体做法地我在源码里用了一个“结构导出”工具遍历剪枝后的net输出一个新的yaml包括每个层的通道数、RepCSP等模块的重复次数。这个工具写起来比较繁琐但倒是很实用我后面有空会单独写一篇源码解析。在最简单的情形下如果你只是要推理和导出就保存整个model.state_dict()然后改造模型定义文件来适配新的通道数即可。我实际项目里是用“读取模型结构然后修改输出通道”的方法可以参考torch_pruning官方对ResNet的剪枝后导出代码的思路把对应模块的输入输出通道改掉。3.2 精度恢复与推理加速验证微调之后需要做两件事验证精度和验证速度。精度验证可以用ultralytics的val命令或者自己写脚本。我自己习惯统计mAP50和mAP50-95两个指标。还是拿前面那个交通标志检测的任务举例指标原始yolov8s剪枝后不微调剪枝后微调50轮mAP500.8220.7370.806mAP50-950.5910.5020.575参数总量11.2M7.8M7.8MCPU推理耗时640x64065ms42ms42msGPU推理耗时RTX306011.2ms7.1ms7.1ms可以看到剪枝后参数减少约30%推理速度提升明显。但要是不微调mAP50掉的4.3个百分点确实很难看。微调后只掉1.6个点达到了业务预期。推理加速我是用下面的脚本测的import time import torch from ultralytics import YOLO def benchmark(model_path, img_size640, warmup10, runs50): model YOLO(model_path) inputs torch.rand(1, 3, img_size, img_size) for _ in range(warmup): model.predict(inputs, imgszimg_size, verboseFalse) torch.cuda.synchronize() start time.time() for _ in range(runs): model.predict(inputs, imgszimg_size, verboseFalse) torch.cuda.synchronize() avg (time.time() - start) / runs return avg注意直接用YOLO类做推理会有很多预处理后处理的开销如果想更精确地测模型网络本身的速度可以取model.model模块直接跑forward。实际部署时我们还需要把前后处理部分优化掉才能拿到真实帧率。3.3 导出ONNX与TensorRT部署剪枝微调完最终还是要部署的。我这边常用的是导出ONNX然后转TensorRT。新版本的ultralytics直接支持exportfrom ultralytics import YOLO model YOLO(yolov8s_pruned.yaml) model.load(yolov8s_pruned_finetuned.pt) model.export(formatonnx, opset12, simplifyTrue)然后ONNX转TensorRT引擎trtexec --onnxyolov8s_pruned_finetuned.onnx \ --saveEngineyolov8s_pruned_finetuned.engine \ --fp16 \ --minShapesimages:1x3x640x640 \ --optShapesimages:1x3x640x640 \ --maxShapesimages:1x3x640x640我在TensorRT上测出来的速度比PyTorch里快了近一倍。剪枝后模型本身计算量小了加上TensorRT的层融合边缘设备上跑起来很舒服。但这里有一个很隐蔽的坑剪枝后模型结构改变了导出ONNX时如果用了不当的dynamic_axes配置或者某个层因为通道数变化导致名字变化可能会导出失败。我的建议是先导出静态形状的ONNX跑通后再考虑动态形状。因为剪枝后模型的Channel已经固定了一般动态batch就够了宽高动态反而容易出问题。4. 常见问题与踩坑记录4.1 剪枝后层维度不匹配怎么办这是最容易踩的坑。尤其当你手动剪掉某个卷积的输出通道后该卷积的下一层比如cat、add操作维度对不上直接报RuntimeError。我遇到的最典型的情况是C2f模块内部的拼接。C2f结构会把经过多个bottleneck分支的结果和一个skip连接拼在一起torch_pruning虽然能识别cat但有时候因为模块太嵌套会漏掉某个分支。这时候需要自己干预。排查方法把剪枝前后的模型分别用随机输入跑一遍forward用hook检查所有层的输出形状定位是哪个层开始对不上的。我写过一个辅助函数def check_shapes(model, input_tensor): shapes {} def hook_fn(name): def hook(module, inputs, outputs): shapes[name] outputs.shape return hook hooks [] for name, module in model.named_modules(): if isinstance(module, (torch.nn.Conv2d, torch.nn.BatchNorm2d)): hook module.register_forward_hook(hook_fn(name)) hooks.append(hook) model.eval() with torch.no_grad(): model(input_tensor) for h in hooks: h.remove() return shapes运行后对比剪枝前后的层输出形状差异你很快就能找出问题点。解决办法一般有两个一是把问题层也纳入剪枝计划保证它和前面的层同步剪二是对特殊模块手工修改结构比如调整C2f中某个通道数。我的经验是千万不要试图跳过依赖分析手动去改某个层的通道数99%会漏。务必用DepGraph自动生成剪枝计划。4.2 剪枝后精度暴跌的原因分析有时候剪枝比例明明不大精度却掉得很夸张。我见过最离谱的一次是15%的剪枝比例mAP直接从0.7掉到0.3。反复排查后发现问题出在剪枝通道选择错了——BN gamma最小的通道并不等价于最不重要的通道。为什么会这样因为YOLOv8s在训练时使用了强数据增强和EMA指数移动平均更新权重。EMA的模型权重和保存出来的BN统计量其实和当前模型并不完全一致。如果你的训练过程还用了混合精度BN的gamma分布可能不那么“干净”。此外如果数据集类别较少或者某些通道对特定类别特别重要那么全局按gamma排序会误伤关键通道。针对这个情况我后来改进了策略不再只按gamma绝对值排序而是结合一个“敏感度权重”在计算重要性时乘上该通道后续连接的Conv1x1对应权重的范数。虽然源码复杂了一些但精度稳定了很多。如果不想搞太复杂至少可以在剪枝前先跑一次bn_calibration也就是用几个batch数据重新统计BN的mean/var和gamma的分布这样剪得更准。具体操作def calibrate_bn(model, dataloader, num_batches50): # 比如用train集中50个batch重新估计BN参数 model.train() for i, (imgs, labels) in enumerate(dataloader): if i num_batches: break with torch.no_grad(): model(imgs) model.eval()强制model.train()会更新BN的running_mean/running_var同时batch统计量更新gamma的分布。实测这个操作能让剪枝后精度再涨0.5到1个点。4.3 源码级调试经验最后分享几个调试心得。第一torch_pruning的DepGraph在构建时要求example_inputs的形状必须是有效的YOLOv8s的model需要在train模式下跑一遍吗不要用eval。而且输入不要太多通道正常1x3x640x640就好。第二如果你剪的是Backbone或Neck但模型有检测头依赖这些层输出的特征图那么在剪枝计划里一定要确保检测头对应的输入通道也被同步更新。torch_pruning其实会自动处理但如果你的检测头用的是自定义插值或concat可能不会自动。我调试时就在Detect的forward里发现过特征图通道对不上最后不得不把Detect内的相关卷积也纳入剪枝或者是用1x1卷积把通道数对齐回去。第三剪枝后的模型微调前一定要先固定随机种子。因为剪枝已经改变了模型结构如果训练配置变了精度波动会很大。我在源码里加了一行def set_seed(seed42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed)第四建议把剪枝前后的模型参数量、FLOPs算一下好量化收益。ultralytics里自带profile方法也可以用thop库from thop import profile, clever_format flops, params profile(net, inputs(torch.zeros(1,3,640,640),), verboseFalse) flops, params clever_format([flops, params], %.3f) print(f剪枝模型: FLOPs{flops}, Params{params})4.4 剪枝源码的扩展方向我现在这套yolov8s剪枝源码本质上不止适用于YOLOv8s稍微调整一下过滤规则也能用到YOLOv5、YOLOX这些模型上。核心逻辑都是一样的找可剪层、算通道重要性、构建依赖图、执行剪枝、重新导出结构。后续我还打算在源码里加上自动搜索最佳剪枝比例的功能用类似二分法或贝叶斯优化跑一轮agent找出在精度约束下的最大剪枝率。如果大家对这个感兴趣我可以把源码整理一下放出来。我个人在实际操作中的体会是模型剪枝拼的不是算法复杂度而是对模型结构和训练细节的熟悉程度。你把YOLOv8s结构吃透了剪枝就是很自然的事情。踩过几次坑之后我现在剪枝基本能一次到点不再像最初那样动不动维度报错、精度崩盘了。希望这篇文章能帮你少走弯路。
返回列表