我手头有一个训练好的yolov8s模型,跑在Jetson Orin上,单帧推理大概20多毫秒,还想往15毫秒以内压;显存也吃紧,batch稍微拉大一点就OOM。问了一圈,方案无非是剪枝、量化、蒸馏。量化动静大,蒸馏费时间,最后我选了yolov8s剪枝这条路线,而且坚持要拿到能自己改的源码,不能只靠黑盒工具点几下。这篇文章就把我跑通的yolov8s模型剪枝源码和完整流程整理出来,把原理、关键代码、参数设置、踩坑记录一次性说清楚。适合手里已经有yolov8s权重、想把它部署到边缘设备或者低显存环境的人;如果你只是对剪枝算法本身感兴趣,从第1节和第2节同样能拿到足够干的货。
1. 剪枝到底在剪什么:原理与方案选型
1.1 神经网络为什么能剪
深度学习模型的参数冗余是常态。一个卷积层里几十上百个卷积核,真正对最终结果起决定性作用的可能只有一部分,其余的贡献趋近于零。这不是玄学,你去看训练收敛后卷积层的权重分布,大量数值集中在0附近,说明这些通道基本处于“半退休”状态,留着它们只是占算力和显存。
YOLOv8s的结构比v5更规整,backbone用C2f模块替换了C3,neck还是PAN-FPN,head换成了Decoupled Head(解耦头)。但规整不等于没有冗余。C2f里堆叠的Bottleneck反复提取特征,中间存在大量相似表达;PAN-FPN在融合多尺度特征时,很多通道是重复计算。你可以这样理解:一个公司里真正扛事的关键员工可能只占三成,其余是辅助、备份、甚至摸鱼的,剪枝就是把这些“不产生实际价值的编制”裁掉,公司和剪枝后网络的共同点是——活儿还得照常干。
剪枝的本质不是删除神经元的“概念”,而是删除具体的计算单元。在YOLOv8s这种CNN模型里,最常剪的是卷积层的输出通道(也就是该层卷积核的个数),一旦通道数减少,后面所有依赖这个输出的层都得跟着改,这就带出了依赖图的概念,我在第2.3节会详细讲。
1.2 结构化剪枝和非结构化剪枝怎么选
剪枝分两大流派:非结构化剪枝和结构化剪枝。网上讨论“剪枝算法”的时候经常把两者混在一起说,实际工程中差别非常大。
非结构化剪枝是直接把权重矩阵里绝对值很小的单个参数置零,得到一个稀疏矩阵。它不改变张量形状,精度保留通常也更好,但问题在于:通用硬件不认稀疏矩阵。GPU、NPU上如果没有配套的稀疏卷积算子,剪完的模型跑起来不仅不快,反而因为要额外处理索引而更慢。Jetson这类边缘设备上基本没有成熟的高效稀疏算子库,所以非结构化剪枝更适合学术研究和特殊芯片场景。
结构化剪枝是整条通道、整个卷积核地删。删完张量形状直接变小,下一层的输入通道数也跟着变。它不需要专门算子支持,任何深度学习框架、任何推理引擎都能直接吃,部署友好性极高。代价是精度掉得相对快一些,需要靠微调拉回来。
我最终选了结构化通道剪枝,就是看中它“剪完直接跑”的优点。另外还要区分预剪枝和后剪枝:预剪枝在训练前就对随机初始化的网络动手,剪完再训练,这种方式几乎没人用,因为随机初始化下没法判断哪些通道重要;后剪枝先训练一个完整baseline,剪完再微调,也就是“先学后砍再补课”,目前主流方案都是在后剪枝或者训练中做稀疏化。
1.3 重要性评估算法:为什么我用BN缩放因子
剪枝算法里最核心的问题只有一个:怎么判断哪些通道该剪。学术界给过一堆答案,我在实际工程里接触到的无非这几种:
| 评估方式 | 核心思路 | 工程难度 | 实战效果 |
|---|---|---|---|
| 权重幅度 | 看卷积核权重绝对值之和 | 低 | 能用,但容易忽略深层语义 |
| BN缩放因子 | 看BN层γ参数幅度 | 低 | 稳定,适合带BN的网络 |
| 泰勒展开 | 用梯度×激活近似损失变化 | 高 | 理论严谨,实现复杂 |
| Hessian信息 | 二阶导数评估重要性 | 很高 | 效果上限高,算力消耗大 |
对YOLOv8s来说,最实用的是BN缩放因子方案。原因很直接:YOLOv8s几乎每个Conv模块都是“Conv2d + BatchNorm2d + SiLU”的组合,BN层是现成的,每个通道都有一个可学习的缩放因子γ。训练过程中,如果给γ加上L1正则,冗余通道的γ会被压向0甚至变成负值,γ绝对值越小的通道,对后续特征图的贡献越弱,剪掉它影响最小。这套思路源自Network Slimming那篇论文,现在大量工程实现都是它的变体。
我用的完整技术路线是:训练中稀疏化(给BN的γ加L1正则)→ 统计γ分布并设定阈值 → 结构化通道剪枝 → 微调恢复精度。这套流程对yolov8s模型来说性价比最高,下面我拆开讲源码和细节。
2. yolov8s剪枝源码核心拆解
2.1 把BN层γ参数当成“法官”
YOLOv8s里BN层的γ参数就是剪枝判断的依据。每个通道对应一个γ,训练时γ会跟着梯度更新,加上L1正则后,不重要的通道γ会被压向0。剪枝前第一件事就是把所有BN层的γ收集起来,看它们的分布。
统计代码很直接,遍历模型里的所有BatchNorm2d模块就行。有一个细节必须注意:Detect头里其实也有BN,但Detect头最后一层卷积的输出通道受类别数和DFL回归参数约束,不能参与剪枝,所以统计时要排除整个detect模块名。
import torch def collect_bn_gamma(model, exclude_names=("detect",)): gamma_list = [] names = [] for name, module in model.named_modules(): if isinstance(module, torch.nn.BatchNorm2d): if any(ex in name for ex in exclude_names): continue g = module.weight.data.detach().cpu().abs() gamma_list.append(g.view(-1)) names.append((name, g.numel())) all_gamma = torch.cat(gamma_list) return all_gamma, names收集完γ以后,把它画成直方图或者打印统计值。我一般看min、median、p95三档:如果min已经非常接近0,而median还比较大,说明稀疏化训练起了作用,冗余通道明显;如果min还是正的且离0很远,说明稀疏化强度不够,回去调λ。这里有一个很关键的点:我收集的是abs()绝对值,因为稀疏化训练后有一部分γ会变成负的,绝对值才能反映“通道输出强度”。
另外提醒一下,别把输入层第一个卷积的BN排除在外。虽然第一层卷积剪了会影响输入通道数,但它的输出通道是可以剪的,而且因为第一层特征图分辨率最大,剪掉它的收益非常明显。
2.2 全局阈值还是分层阈值:剪枝率的门道
拿到全部γ值之后,就要决定剪多少、按什么标准剪。这里有两条路线:全局阈值和分层阈值。
全局阈值就是把所有BN层的γ合并成一个列表,按目标剪枝率直接取分位数作为阈值。比如我想剪掉40%的通道,就取γ分布的第40百分位数作为threshold,所有γ小于threshold的通道全部剪掉。优点是整体剪枝率精确可控,缺点是“一刀切”——浅层的γ普遍较小,可能被剪到只剩几个通道,直接导致特征提取能力崩溃。
分层阈值是每层单独设置保留比例,比如每层都只保留60%的通道。优点是分布均匀,不会出现某一层被剪秃的情况;缺点是不区分层的重要性,对某些冗余本就少的层造成了不必要的损伤。
我的实践证明,最稳的是“全局阈值+最小保留通道数限制”的组合方案:
min_channels = 16 # 每层至少保留的通道数 def compute_threshold(all_gamma, prune_ratio): return torch.quantile(all_gamma, prune_ratio) def get_pruning_idxs(model, threshold, min_channels=16): pruning_idxs = {} for name, module in model.named_modules(): if isinstance(module, torch.nn.BatchNorm2d): if "detect" in name: continue num_channels = module.weight.data.shape[0] idxs = (module.weight.data.abs() < threshold).nonzero().view(-1) if (num_channels - len(idxs)) < min_channels: # 如果剪完低于下限,调整保留数量 keep = torch.argsort(module.weight.data.abs())[-min_channels:] idxs = torch.tensor([i for i in range(num_channels) if i not in keep]) pruning_idxs[module] = idxs.tolist() return pruning_idxs这里min_channels=16是我踩过几次坑之后定的经验值。浅层卷积分辨率高、参数量大,但通道数本身不算多,剪太狠容易让后面的层拿不到有效特征。加了这个下限之后,即使全局阈值选得激进一点,也不至于把某一层剪空。
剪枝率怎么定?我建议第一次别贪心。先用10%、20%、30%、40%四档做一次“剪枝率敏感性测试”:每个比例剪完,微调10个epoch,看mAP掉多少。很多时候你会发现从20%到30%是质变点,掉点突然变多,那说明网络真正冗余的部分大概就在这个区间附近。这个测试总共花不了多长时间,比凭感觉定一个50%然后剪完发现模型废掉要划算得多。
2.3 通道联动:为什么不能只剪一个卷积
这是剪枝源码里最容易翻车的地方。你以为剪的是“一个卷积层的几个通道”,实际上卷积和卷积之间存在复杂的依赖关系:A卷积输出作为B卷积的输入,那么A剪掉的通道,B的输入通道必须同步调整;在PAN-FPN的cat节点,两个输入分支的通道数必须一致;add节点更严格,形状要完全对齐。
手写这个联动逻辑极其痛苦。YOLOv8s有SPPF、C2f、PAN-FPN、上采样、cat、add,各种模块交织在一起,一个通道索引算错,后面全乱。所以我自己写剪枝源码时选择了基于依赖图的方案——torch_pruning这个库,它能自动分析模型里所有张量的上下游关系,生成一棵依赖树。我剪一个BN层,它会自动找到对应的卷积、后续的卷积、以及所有经过cat和add的兄弟路径,一次性生成完整的剪枝计划并执行。
核心剪枝代码大概长这样:
import torch_pruning as tp model.eval() example_inputs = torch.randn(1, 3, 640, 640).to(device) dg = tp.DependencyGraph().build_dependency(model, example_inputs=example_inputs) threshold = compute_threshold(all_gamma, 0.3) for module, idxs in get_pruning_idxs(model, threshold).items(): if len(idxs) == 0: continue plan = dg.get_pruning_plan(module, tp.prune_batchnorm, idxs=idxs) print(plan) plan.exec()build_dependency这一步会跑一次前向,让库把每一层的输入输出张量形状记录下来,进而构建依赖关系。get_pruning_plan返回的就是一个完整的剪枝执行计划,它会显示“当前剪这个通道,连带需要调整哪些层的输入输出通道”。执行前务必打印这个plan看一眼,确认没有把Detect头里的输出层卷进去。
Detect头的处理是yolov8系列剪枝最重要的一个坑。YOLOv8的head是解耦头,分类分支和回归分支各自有几层卷积,而回归分支还接了一个DFL模块,输出通道数是4 * reg_max(reg_max默认16),这个维度是硬约束,不能因为剪枝而改变。实操里我直接把detect模块整体豁免剪枝,虽然这会丢失一部分压缩空间,但换来了稳定性和省心。后面如果还想榨性能,再考虑对detect内部除输出层以外的卷积做精细剪枝,但那就需要逐层验证了。
3. yolov8s剪枝实操全流程
3.1 环境和baseline准备
剪枝前先把基础打好。我的环境是Python 3.10、PyTorch 2.1、ultralytics 8.2.x、torch_pruning 1.4.x,这几个版本组合下来没有遇到明显兼容问题。注意torch_pruning的API在不同版本之间有过变化,如果build_dependency的参数形式不一致,以你安装版本的官方文档为准。
第一步是跑基线。如果下游任务用COCO预训练权重直接部署,当然可以省去baseline训练,但如果你有自定义数据集,强烈建议先用自己的数据微调或者重新训练一个yolov8s出来,再在这个基础上剪枝。原因很简单:COCO的分布和你的业务分布大概率不一样,在COCO上冗余的通道在你的数据上未必冗余。
训练baseline时有一个原则:固定随机种子。剪枝前前后后要对比mAP,如果每次训练数据划分、增强顺序都不一样,mAP的波动会被你误判成剪枝的损失。
yolo detect train data=your_dataset.yaml model=yolov8s.yaml epochs=100 imgsz=640 seed=42训练完记录4个数字:mAP50-95、参数量、FLOPs、单帧GPU延迟。剪枝完成后这4个数就是你的对照基线。这几个数后面全部要用,少记一个都得重跑。
3.2 稀疏化训练:让γ学会“自我淘汰”
稀疏化训练是整个流程中最容易做错的一步。目的是让BN层的γ尽量往0压,同时不破坏模型原有精度。做法是在原有损失函数后面加上一项L1正则:loss = base_loss + λ * sum(|γ|)。
λ的取值很微妙。太小了压不动γ,训练完跟没加一样;太大了模型忙着收敛γ,正常检测任务精度崩掉。我的经验是从1e-4开始试,观察γ的min值和mAP的变化,如果训练了20个epochγ的min还在0.1以上,调大到1e-3再试一次。
在ultralytics里集成稀疏化损失,我用的方式是自定义一个DetectionModel子类,在forward返回的loss字典里追加一项:
import torch import torch.nn as nn from ultralytics.nn.tasks import DetectionModel class SparseDetectionModel(DetectionModel): def __init__(self, cfg="yolov8s.yaml", ch=3, nc=None, verbose=True, lam=1e-4): super().__init__(cfg, ch, nc, verbose) self.lam = lam def forward(self, x, *args, **kwargs): loss, outputs = super().forward(x, *args, **kwargs) if self.training and isinstance(loss, dict): l1 = 0.0 for m in self.modules(): if isinstance(m, nn.BatchNorm2d): l1 += torch.abs(m.weight).sum() loss["l1"] = self.lam * l1 return loss, outputs然后训练时把模型替换进去:
from ultralytics import YOLO model = YOLO("yolov8s.yaml") model.model = SparseDetectionModel("yolov8s.yaml", nc=num_classes, lam=1e-4) model.train(data="your_dataset.yaml", epochs=150, imgsz=640, seed=42)我不建议整个训练过程全程加稀疏正则。通常的做法是:正常训练baseline到mAP稳定 → 在最后三分之一epoch里打开稀疏损失 → 关掉稀疏损失再正常训练10个epoch让γ收敛。后一步很多人忽略,它很重要,因为L1正则会把γ往负方向推,直接剪枝时某些γ还是负的,剪完再微调恢复起来更吃力,给一点正常训练时间让分布稳定下来,剪出来的通道结构才靠谱。
另外一个很有用的经验:对BN层的参数关闭weight decay。PyTorch的优化器默认会对所有参数加weight decay,BN的γ本身很小,如果weight decay也在压它,会和L1正则的效果混在一起,导致γ分布不稳定。我给BN单独设一个参数组来处理:
bn_params = [] other_params = [] for name, param in model.named_parameters(): if "bn.weight" in name: bn_params.append(param) else: other_params.append(param) optimizer = torch.optim.SGD( [ {"params": other_params, "weight_decay": 1e-4}, {"params": bn_params, "weight_decay": 0.0}, ], lr=0.01, momentum=0.937, )3.3 执行剪枝:备份、剪枝、保存结构
稀疏化训练完成后,进入剪枝阶段。动手之前做两件事:备份原始完整权重,记下当前随机种子。前者是为了剪坏了能随时回到原点,后者是为了让剪枝操作本身可复现。
然后把模型切到eval模式,构建依赖图,按第2.2节的方法计算阈值并执行剪枝。执行完以后一定要做三件事:打印每层剪枝前后的shape变化、统计删除了多少通道、跑一次前向确认模型能正常输出。
pruned_model = prune_model(model.cpu(), torch.randn(1, 3, 640, 640)) total_before = sum(m.weight.data.shape[0] for m in pruned_model.modules() if isinstance(m, nn.Conv2d)) total_after = sum(m.weight.data.shape[0] for m in pruned_model.modules() if isinstance(m, nn.Conv2d)) print(f"C2f/Conv通道数: {total_before} -> {total_after}")剪枝后保存模型对象,不要再只存state_dict。这里踩过一个很深的坑:剪枝改变的是模型结构本身,之后再加载时如果只是load_state_dict,类名、通道数都对不上。最省心的做法是直接把整个模型对象保存下来:
torch.save(pruned_model, "yolov8s_pruned_full.pt")如果非要走ultralytics的YOLO加载流程,你还需要同步保存一份经过修改的模型结构说明文件,否则YOLO("xxx.pt")加载时按原始yaml重建结构,会直接报shape mismatch。我后面第4.2节会再讲这个坑。
3.4 微调:把精度拉回来
剪完的模型精度必然掉,微调就是把掉出来的精度一点点补回来。这一阶段我试过两种策略,各有取舍。
一种是直接从头微调:把剪完的模型当作一个“新的随机初始化模型”,用cosine学习率调度,从lr=0.01左右重新训练。这种方式收敛快,适合剪枝比例不高(比如30%以内)的情况。另一种是两阶段微调:先冻结backbone,只让neck和head训练20个epoch,再解冻全部层微调。这种方式适合剪枝比例较高、模型损伤较大的情况,防止一上来全网络一起动导致训练不稳定。
微调epoch一般是原始训练的一半到等量。比如原始训练100个epoch,微调就设置50到100个epoch。如果100个epoch还拉不回精度,多半不是epoch不够,而是剪多了。
微调期间还要调整数据增强强度。ultralytics默认开了Mosaic马赛克增强,这在训练初期帮助很大,但在微调阶段高强度的Mosaic会让模型难以收敛回原来的分布。我一般设置mosaic=0.5左右,或者在前一半epoch用Mosaic、后一半关掉。HSV增强、平移、缩放这些轻量增强保持默认,给模型一点正则化,防止微调阶段过拟合到训练集。
微调完成后,对比剪枝前后的mAP和参数量。一个健康的剪枝结果应该是:参数和FLOPs下降30%到50%,mAP掉点控制在1到2个点以内,甚至经过充分微调反超baseline也不是没可能。
3.5 部署验证:导出ONNX和TensorRT
剪枝和微调都完成之后,最后一步是部署验证。不管目标平台是Jetson、PC还是手机,我都会先导出一版ONNX,确认结构没问题。
yolo export model=yolov8s_pruned_full.pt format=onnx opset=12 imgsz=640导出后检查输出张量的shape,YOLOv8检测头的输出维度必须和类别数、DFL参数保持一致。如果这里发现shape不对,基本可以断定是剪枝时Detect头被波及了,回到第2.3节的豁免逻辑去查。
TensorRT推理还要重新build engine:
trtexec --onnx=yolov8s_pruned_full.onnx --fp16 --saveEngine=yolov8s_pruned_full.engine剪枝后的模型必须重新构建engine,不能拿之前的engine硬跑。构建完成后对比剪枝前后的实测延迟和显存占用。这一轮下来你会得到一个非常直观的收益:模型体积变小、显存占用下降、延迟缩短。如果你的目标本身就是低显存运行模型,剪枝后把batch size拉大一倍通常问题不大。
4. 常见问题与排查技巧实录
4.1 剪完mAP掉太多,怎么判断是哪里出了问题
这是我最常被问的问题。先把尺度对齐:mAP50-95掉3个点以内,通过微调完全可以接受;掉5个点以上,说明剪枝策略有问题。
排查顺序建议从这几个方面走:剪枝比例是不是一步到位了?我之前提过,如果目标是50%,最好分两次:先剪30%,微调,再剪剩下的部分。一步到位50%对yolov8s来说掉点会非常明显。
再看全局阈值是不是把某些敏感层剪秃了。打印剪枝后每层的通道数分布,如果某个层只剩个位数通道,这一层基本就废了。解决办法就是第2.2节说的,设置每层最小通道数下限,把敏感层加进豁免名单。
还有微调策略的问题。我见过很多人剪完直接用原始训练配置再跑一遍,学习率、数据增强都不改,结果mAP死活回不来。微调阶段的数据增强强度一定要降,学习率调度一定要重新设置,这两点做不到位,模型很难收敛回原来的状态。
4.2 加载剪枝模型报shape mismatch
这个报错基本等于告诉你:模型结构对不上。绝大多数原因是只保存了state_dict,然后试图用ultralytics的YOLO类去加载,ultralytics加载时会根据内置的yaml描述重建模型结构,而原始yaml里的通道数是没剪之前的,自然对不上。
我的建议是剪枝后保存完整模型对象,加载时用torch.load配合map_location恢复模型,然后再转成推理或训练模式。如果你确实需要导出成ultralytics能直接加载的模型文件,那就必须同时维护一份与剪枝后通道数匹配的模型结构描述文件,并且逐层核对。这个工作量不小,我建议只在最终交付时做,调试阶段都用完整模型对象。
4.3 稀疏化训练跑完,γ几乎没变化
这种情况多半是λ太小或者BN参数没被优化器更新。先打印γ的min、median、p95看看,如果min一直是正数而且离0很远,说明L1的梯度压根没起作用。
检查优化器参数分组是不是把BN参数漏掉了,如果BN的γ不在任何一个param_group里,它根本不会被更新。另一个容易被忽略的是weight decay,BN层的γ很小,如果weight decay也在压它,L1正则那点梯度会被抵消掉,出现“加了和没加一样”的现象。我一般会把BN参数的weight decay显式设成0,再单独用λ来控制γ的稀疏度。调大一个数量级试试也是个快速验证方法,λ从1e-4提到1e-3,跑10个epoch看γ分布有没有明显下沉。
4.4 剪完FLOPs降了,推理延迟却没变
这个问题的本质是:FLOPs是理论计算量,GPU的实际耗时还受内存带宽、算子调度、kernel launch开销影响。剪掉一些计算量大的卷积,但如果瓶颈在SPPF里的MaxPool、上采样层或者后续的NMS后处理,整体延迟当然不会有明显变化。
排查方法是用profiling工具看每一层的耗时。TensorRT自带的trtexec --dumpProfile就能输出每层耗时排名,找出真正吃时间的是哪些层。如果瓶颈在上采样和MaxPool这种不可剪的层,通道剪枝对延迟的改善就是有限的,这时候要么改输入分辨率,要么走量化和算子融合路线。另外别忘了:剪枝后一定要重新build TensorRT engine,用旧的engine测速等于白剪。
| 常见问题 | 现象 | 主要原因 | 解决思路 |
|---|---|---|---|
| mAP掉点多 | 掉5个点以上 | 剪枝率过高、敏感层被剪 | 降低比例、分步剪、设置通道下限 |
| 加载报错 | shape mismatch | 只存权重、结构未同步 | 保存完整模型对象或同步结构文件 |
| 稀疏训练无效 | γ分布无变化 | λ太小、BN未更新 | 检查参数组、关BN的weight decay |
| FLOPs降但延迟未降 | 推理速度没变化 | 瓶颈在不可剪层 | 用profiler定位,考虑量化融合 |
剪枝这件事真正花时间的从来不是剪的动作本身,半小时就能把源码跑通。真正决定成败的是前面的稀疏化训练和后面的微调,这两步做扎实了,剪枝就是水到渠成的事。最后分享一个我很依赖的小习惯:每次剪枝前都固定随机种子,剪完立刻导出一版ONNX,用同一张输入图对比剪枝前后中间层特征图的余弦相似度,哪一层被剪坏了马上就能看出来。希望这篇yolov8s剪枝源码的实战记录对你有用,这套流程换到其他带BN的检测模型上,思路也完全通用。