AI 剪枝踩坑记录
项目本身不复杂:一个文本分类任务,在服务器上训练好的 BERT 模型,分类准确率能达到 92% 以上,但部署时问题来了。
GPU 显存只有 8GB,要塞的模型却动不动就几 GB,压缩过程就是在一堆"这个参数能不能删"的纠结里推进的。
问题背景
项目本身不复杂:一个文本分类任务,在服务器上训练好的 BERT 模型,分类准确率能达到 92% 以上,但部署时问题来了。目标设备是某款边缘计算盒子,显存 8GB,系统和其他服务已经占了 2GB 左右,剩下的空间要塞模型、推理框架和一些预处理模块。
初始模型参数量约 110M,按 fp16 存储就要 200MB 左右。加上推理引擎的内存开销,理论上能塞进去,但实际跑起来发现内存占用经常超过预期,而且推理延迟在 40-50ms 之间,无法满足实时性要求。
需要把模型压缩到约 60-70M 参数,同时保证分类准确率下降不超过 2-3 个百分点。看起来压缩目标不算激进,但实际操作起来,剪枝、量化、蒸馏这些手段都要试一遍才知道哪个组合最管用。
剪枝方案选型
剪枝大体分成两类:结构化剪枝和非结构化剪枝。
结构化剪枝剪的是整张"纸"——删掉整个卷积核、整个神经元或者整个通道。好处是剪完之后模型结构还是规整的,推理引擎不需要特殊优化就能加速;坏处是搜索空间小,精度损失比较大。
非结构化剪枝剪的是"纸上的墨点"——删掉具体的某个权重参数。好处是可以更精细地控制哪些参数该删,精度损失小;坏处是剪完后模型会变得稀疏,推理引擎需要专门的稀疏计算库才能利用到加速,否则稀疏的好处全被内存访问的额外开销吃掉了。
从算力条件和工具链考虑,我们打算先试结构化剪枝,看看精度损失能不能接受。如果精度掉得太多,再考虑非结构化剪枝加上稀疏推理库的方案。
结构化剪枝实践
剪枝思路很简单:训练时先用梯度信息评估每个参数的重要程度,然后把不重要的参数删掉,再微调一段时间恢复精度。
但实际操作有几个关键点:
1. 重要性评估方式
用最朴素的 L1/L2 正则化方法,给每个参数加上一个基于权重大小的"重要性分数"。权值绝对值越大的参数越重要,越小的越可能是冗余的。
import torch
import torch.nn as nn
def calculate_importance(model):
importance = {}
for name, param in model.named_parameters():
if 'weight' in name:
# 使用 L1 范数作为重要性指标
importance[name] = torch.norm(param.data, p=1, dim=tuple(range(1, param.dim())))
return importance
这个方法简单直接,但有个问题:有些权重虽然绝对值小,但对特定输入很重要;有些权重虽然大,但长期以来几乎不起作用。所以后来又试了基于梯度的一阶重要性评估:
def calculate_gradient_importance(model, dataloader):
importance = {}
model.eval()
for name, param in model.named_parameters():
if 'weight' in name:
importance[name] = torch.zeros_like(param.data)
for batch in dataloader:
outputs = model(batch)
loss = criterion(outputs, batch.labels)
loss.backward()
with torch.no_grad():
for name, param in model.named_parameters():
if 'weight' in name:
# 累积梯度绝对值
importance[name] += torch.abs(param.grad * param.data)
return importance
基于梯度的方法计算开销更大,但评估准确性确实提升不少。
2. 剪枝策略
确定了重要性后,就是怎么剪的问题。这里有几个常见策略:
- 全局剪枝:不按层级,把所有参数排个队,直接按重要性砍掉最不重要的那些。这种剪枝比较激进,容易导致某些层被完全砍空。
- 局部剪枝:每一层单独评估,按固定比例剪掉每层最不重要的参数。这种剪枝比较保守,但能保证每层都有保留一些信息。
- 渐进式剪枝:分多次剪,每次只剪一点,中间穿插微调。这种剪枝最耗时,但精度损失最小。
我们采用了渐进式局部剪枝的策略,原因是:
- 全局剪枝容易砍掉某些层的关键信息,导致模型崩塌;
- 一次性剪枝太多,微调阶段很难恢复;
- 虽然耗时,但我们的环境允许长时间训练。
def progressive_pruning(model, train_loader, test_loader, target_sparsity=0.6, num_iterations=10):
current_sparsity = 0.0
sparsity_increment = target_sparsity / num_iterations
for iteration in range(num_iterations):
# 计算当前重要性
importance = calculate_gradient_importance(model, train_loader)
# 计算每层的剪枝阈值
thresholds = {}
for name, imp in importance.items():
if 'weight' in name:
# 当前轮次的目标稀疏度
target = min(current_sparsity + sparsity_increment, target_sparsity)
# 计算阈值:剪掉最不重要的 target 比例的参数
flat_imp = imp.flatten()
k = int(len(flat_imp) * target)
thresholds[name] = torch.kthvalue(flat_imp, k + 1).values.item()
# 执行剪枝
for name, param in model.named_parameters():
if 'weight' in name and name in thresholds:
mask = torch.abs(param.data) >= thresholds[name]
param.data *= mask.float()
# 微调恢复精度
fine_tune(model, train_loader, epochs=5)
# 评估
accuracy = evaluate(model, test_loader)
print(f"Iteration {iteration + 1}: Sparsity {current_sparsity:.2f}, Accuracy {accuracy:.2%}")
current_sparsity += sparsity_increment
return model
3. 剪枝后的微调
剪枝会破坏模型的结构,微调阶段是恢复精度的关键。微调时需要注意:
- 学习率要小:剪枝后的模型对参数变化更敏感,学习率太大容易震荡。
- 训练时间要够:虽然只是微调,但往往需要比普通训练更长的轮次才能恢复精度。
- 数据要充足:剪枝后的模型更容易过拟合,需要更多的训练数据支撑。
微调阶段我们用了初始学习率 1e-5,比正常训练小了一个数量级,训练了 20 个 epoch,才勉强把精度拉回到可接受范围。
结构化剪枝的坑
结构化剪枝跑下来,精度损失比预期大得多,总结下来有几个坑:
1. 剪枝比例难以控制
理论上剪掉 40% 的参数,模型应该还能保留大部分信息。但实际操作发现,某些层对剪枝极其敏感,稍微剪一点就导致整个层失效,进而影响后续层的表达。
比如 BERT 的 attention 层,剪掉部分头之后,多头机制的优势就没了;再比如 FFN 层,砍掉某些神经元后,整个层的表达能力会急剧下降。
2. 微调阶段不稳定
剪枝后的模型在微调阶段特别容易出现梯度爆炸或梯度消失。因为某些连接被砍掉后,信息传递的路径变窄了,梯度在反向传播时容易积累或消失。
我们试了梯度裁剪、归一化层调整、学习率预热等手段,才勉强稳定下来,但训练时间比预期长了几乎一倍。
3. 部署兼容性问题
剪枝后的模型虽然参数量少了,但结构变了,推理引擎需要重新生成算子代码。某些推理引擎对结构化剪枝的支持不完善,需要手动改模型结构,或者使用特定的 API。
这个过程比想象中麻烦,而且剪枝策略一旦调整,部署流程就要重新适配,增加了工程复杂度。
非结构化剪枝尝试
结构化剪枝效果不如预期,我们开始考虑非结构化剪枝。
非结构化剪枝的核心优势在于:可以删掉任何一个被认为不重要的参数,而不受层结构的限制。这意味着同样的压缩比例下,精度损失会更小。
1. 稀疏张量与稀疏计算
非结构化剪枝的关键在于稀疏张量的表示和计算。一个稀疏张量可以用 CSR(Compressed Sparse Row)或 COO(Coordinate)格式存储:
from torch.sparse import FloatTensor as SparseTensor
# 创建稀疏张量
def create_sparse_tensor(dense_tensor, mask):
indices = torch.nonzero(mask).t()
values = dense_tensor[mask]
sparse_tensor = SparseTensor(indices, values, dense_tensor.size())
return sparse_tensor
但 PyTorch 的稀疏计算支持有限,很多算子还不支持稀疏张量。实际部署时,可能需要专门的稀疏计算库,比如 MKL-DNN 或者 TensorRT 的稀疏计算支持。
2. 非结构化剪枝实现
非结构化剪枝的实现相对简单,基于重要性评估,直接把不重要的参数置零:
def unstructured_pruning(model, importance, sparsity):
for name, param in model.named_parameters():
if 'weight' in name and name in importance:
flat_imp = importance[name].flatten()
# 计算阈值
k = int(len(flat_imp) * sparsity)
threshold = torch.kthvalue(flat_imp, k + 1).values.item()
# 生成掩码
mask = torch.abs(param.data) >= threshold
# 应用掩码
param.data *= mask.float()
return model
这种剪枝方式的好处是灵活度高,可以精确控制哪些参数该删;坏处是推理时需要稀疏计算支持,否则加速效果不明显。
3. 稀疏推理的挑战
稀疏推理的实际加速效果取决于硬件和软件的支持:
- 硬件层面:需要支持稀疏计算的加速卡,比如某些 GPU 的稀疏矩阵乘法指令。
- 软件层面:推理框架需要能识别稀疏张量,并使用优化的稀疏算子。
我们的目标设备硬件支持有限,软件链路也不完善,最后稀疏推理的实际加速效果只有 20% 左右,远不如预期的 2-3 倍加速。
综合方案与结果
折腾了一圈,最后采用了一个综合方案:
- 量化:从 fp32 量化到 int8,模型大小减半,推理速度提升 1.5 倍。
- 非结构化剪枝:剪掉 30% 的参数,精度损失控制在 1% 以内。
- 知识蒸馏:用原始模型作为教师,剪枝后的模型作为学生,通过蒸馏进一步恢复精度。
def distillation_loss(student_outputs, teacher_outputs, labels, temperature=3.0, alpha=0.5):
# 软损失
soft_loss = nn.KLDivLoss(reduction='batchmean')(
F.log_softmax(student_outputs / temperature, dim=1),
F.softmax(teacher_outputs / temperature, dim=1)
) * (temperature ** 2)
# 硬损失
hard_loss = nn.CrossEntropyLoss()(student_outputs, labels)
return alpha * soft_loss + (1 - alpha) * hard_loss
最终结果:

| 指标 | 原始模型 | 压缩后模型 | 变化 |
|---|---|---|---|
| 参数量 | 110M | 60M | -45% |
| 模型大小 | 420MB | 180MB | -57% |
| 推理延迟 | 45ms | 28ms | -38% |
| 分类准确率 | 92.3% | 90.8% | -1.5% |
虽然不是完美结果,但已经满足部署需求,而且精度损失在可接受范围内。
踩坑总结
这次折腾过程有几个关键的踩坑点:
不要迷信理论压缩比:理论上的压缩率往往假设硬件和软件完美支持,实际部署时各种限制会打折扣。
剪枝前要充分评估重要性:基于权重的简单评估容易误判,基于梯度或一阶导数的评估更可靠,但计算开销更大。
微调阶段要给足时间:剪枝破坏了模型结构,恢复需要更长时间的微调,不要急于求成。
硬件限制要提前确认:稀疏计算、量化加速这些特性,不是所有硬件都支持,提前确认能省很多时间。
综合方案往往优于单一方法:剪枝、量化、蒸馏这些手段,组合使用效果更好,但要注意各步骤的顺序和参数调节。
写在最后
模型剪枝不只是技术优化,更像在有限的资源约束下做取舍。有些参数删了可惜,但为了部署只能删;有些优化理论上很好,但工程落地困难只能放弃。
整个过程有点像装修房子:预算有限,空间有限,既要保留核心功能,又要控制在可接受的成本内。最后拿到的方案可能不是最优解,但一定是当前约束条件下的可行解。
这次折腾也让我意识到,边缘端部署和服务器端部署完全是两个世界。服务器端可以堆硬件、堆算力,但边缘端每个字节的内存、每毫秒的延迟都要精打细算。这种约束下的优化,反而更有意思一些。
如果下次再做类似的部署项目,我会更早地考虑硬件和软件的限制条件,把方案设计得更务实一些。毕竟,能跑起来的方案才是好方案。
版权声明: 本文首发于 指尖魔法屋-AI 剪枝踩坑记录(https://blog.thinkmoon.cn/post/277-ai-pruning-structured-unstructured-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。