模型体积大、推理慢,部署到边缘设备总被嫌弃,这是我在跑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(f"Total 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, ratio=0.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_inputs=example_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, idxs=prune_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 gamma:gamma越小越不重要。
- 基于激活值幅度:统计每个通道输出的平均绝对值,越小越不重要。
我的经验是,BN gamma在实际操作中更稳定。因为它已经被模型训练时优化过了,能反映出通道在当前流形中的贡献度。激活值统计需要跑一份数据来统计,很容易受到batch选择的影响,而且检测模型的特征图往往有较多spatial信息,统计量波动很大。
举一个实际例子,我在一个交通标志检测任务里对比过两种策略:
| 策略 | 原始mAP50 | 剪枝后mAP50(30%比例) | 精度损失 |
|---|---|---|---|
| 随机剪 | 0.822 | 0.764 | 5.8% |
| BN gamma剪 | 0.822 | 0.806 | 1.6% |
| 激活值统计剪 | 0.822 | 0.798 | 2.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(data='my_dataset.yaml', epochs=50)不过导出新yaml结构比较麻烦,因为你得知道每一层的out_channels变了多少。我在源码中有个函数:自动遍历剪枝后的net,将各Conv层输出通道变化记录到字典,然后根据原yaml生成新的yaml。这个思路和ultralytics本身的结构解析是兼容的。
具体实现思路:
def generate_pruned_yaml(original_yaml, channel_changes, output_yaml='yolov8s_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轮) |
|---|---|---|---|
| mAP50 | 0.822 | 0.737 | 0.806 |
| mAP50-95 | 0.591 | 0.502 | 0.575 |
| 参数总量 | 11.2M | 7.8M | 7.8M |
| CPU推理耗时(640x640) | 65ms | 42ms | 42ms |
| GPU推理耗时(RTX3060) | 11.2ms | 7.1ms | 7.1ms |
可以看到,剪枝后参数减少约30%,推理速度提升明显。但要是不微调,mAP50掉的4.3个百分点确实很难看。微调后只掉1.6个点,达到了业务预期。
推理加速我是用下面的脚本测的:
import time import torch from ultralytics import YOLO def benchmark(model_path, img_size=640, warmup=10, runs=50): model = YOLO(model_path) inputs = torch.rand(1, 3, img_size, img_size) for _ in range(warmup): model.predict(inputs, imgsz=img_size, verbose=False) torch.cuda.synchronize() start = time.time() for _ in range(runs): model.predict(inputs, imgsz=img_size, verbose=False) torch.cuda.synchronize() avg = (time.time() - start) / runs return avg注意,直接用YOLO类做推理会有很多预处理后处理的开销,如果想更精确地测模型网络本身的速度,可以取model.model模块直接跑forward。实际部署时,我们还需要把前后处理部分优化掉才能拿到真实帧率。
3.3 导出ONNX与TensorRT部署
剪枝微调完,最终还是要部署的。我这边常用的是导出ONNX,然后转TensorRT。新版本的ultralytics直接支持export:
from ultralytics import YOLO model = YOLO('yolov8s_pruned.yaml') model.load('yolov8s_pruned_finetuned.pt') model.export(format='onnx', opset=12, simplify=True)然后ONNX转TensorRT引擎:
trtexec --onnx=yolov8s_pruned_finetuned.onnx \ --saveEngine=yolov8s_pruned_finetuned.engine \ --fp16 \ --minShapes=images:1x3x640x640 \ --optShapes=images:1x3x640x640 \ --maxShapes=images: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_batches=50): # 比如用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(seed=42): 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),), verbose=False) flops, params = clever_format([flops, params], "%.3f") print(f"剪枝模型: FLOPs={flops}, Params={params}")4.4 剪枝源码的扩展方向
我现在这套yolov8s剪枝源码,本质上不止适用于YOLOv8s,稍微调整一下过滤规则,也能用到YOLOv5、YOLOX这些模型上。核心逻辑都是一样的:找可剪层、算通道重要性、构建依赖图、执行剪枝、重新导出结构。
后续我还打算在源码里加上自动搜索最佳剪枝比例的功能,用类似二分法或贝叶斯优化跑一轮agent,找出在精度约束下的最大剪枝率。如果大家对这个感兴趣,我可以把源码整理一下放出来。
我个人在实际操作中的体会是,模型剪枝拼的不是算法复杂度,而是对模型结构和训练细节的熟悉程度。你把YOLOv8s结构吃透了,剪枝就是很自然的事情。踩过几次坑之后,我现在剪枝基本能一次到点,不再像最初那样动不动维度报错、精度崩盘了。希望这篇文章能帮你少走弯路。