AI 剪枝踩坑记录

项目本身不复杂:一个文本分类任务,在服务器上训练好的 BERT 模型,分类准确率能达到 92% 以上,但部署时问题来了。

GPU 显存只有 8GB,要塞的模型却动不动就几 GB,压缩过程就是在一堆"这个参数能不能删"的纠结里推进的。

问题背景

项目本身不复杂:一个文本分类任务,在服务器上训练好的 BERT 模型,分类准确率能达到 92% 以上,但部署时问题来了。目标设备是某款边缘计算盒子,显存 8GB,系统和其他服务已经占了 2GB 左右,剩下的空间要塞模型、推理框架和一些预处理模块。

初始模型参数量约 110M,按 fp16 存储就要 200MB 左右。加上推理引擎的内存开销,理论上能塞进去,但实际跑起来发现内存占用经常超过预期,而且推理延迟在 40-50ms 之间,无法满足实时性要求。

需要把模型压缩到约 60-70M 参数,同时保证分类准确率下降不超过 2-3 个百分点。看起来压缩目标不算激进,但实际操作起来,剪枝、量化、蒸馏这些手段都要试一遍才知道哪个组合最管用。

剪枝方案选型

剪枝大体分成两类:结构化剪枝和非结构化剪枝。

flowchart TD A[神经网络剪枝] --> B[结构化剪枝] A --> C[非结构化剪枝] B --> B1[删除整卷积核/神经元] B --> B2[结构规整 易于部署] B --> B3[精度损失较大] C --> C1[删除单个权重参数] C --> C2[模型稀疏 需特殊支持] C --> C3[精度损失较小]

结构化剪枝剪的是整张"纸"——删掉整个卷积核、整个神经元或者整个通道。好处是剪完之后模型结构还是规整的,推理引擎不需要特殊优化就能加速;坏处是搜索空间小,精度损失比较大。

非结构化剪枝剪的是"纸上的墨点"——删掉具体的某个权重参数。好处是可以更精细地控制哪些参数该删,精度损失小;坏处是剪完后模型会变得稀疏,推理引擎需要专门的稀疏计算库才能利用到加速,否则稀疏的好处全被内存访问的额外开销吃掉了。

从算力条件和工具链考虑,我们打算先试结构化剪枝,看看精度损失能不能接受。如果精度掉得太多,再考虑非结构化剪枝加上稀疏推理库的方案。

结构化剪枝实践

剪枝思路很简单:训练时先用梯度信息评估每个参数的重要程度,然后把不重要的参数删掉,再微调一段时间恢复精度。

但实际操作有几个关键点:

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. 剪枝策略

确定了重要性后,就是怎么剪的问题。这里有几个常见策略:

  • 全局剪枝:不按层级,把所有参数排个队,直接按重要性砍掉最不重要的那些。这种剪枝比较激进,容易导致某些层被完全砍空。
  • 局部剪枝:每一层单独评估,按固定比例剪掉每层最不重要的参数。这种剪枝比较保守,但能保证每层都有保留一些信息。
  • 渐进式剪枝:分多次剪,每次只剪一点,中间穿插微调。这种剪枝最耗时,但精度损失最小。
flowchart LR A[评估参数重要性] --> B[计算剪枝阈值] B --> C[执行剪枝操作] C --> D[微调恢复精度] D --> E[评估模型性能] E --> F{达标?} F -->|否| A F -->|是| G[完成] style A fill:#e3f2fd style B fill:#bbdefb style C fill:#90caf9 style D fill:#64b5f6 style E fill:#42a5f5 style G fill:#2196f3,color:#fff

我们采用了渐进式局部剪枝的策略,原因是:

  1. 全局剪枝容易砍掉某些层的关键信息,导致模型崩塌;
  2. 一次性剪枝太多,微调阶段很难恢复;
  3. 虽然耗时,但我们的环境允许长时间训练。
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 倍加速。

综合方案与结果

flowchart TD A[原始模型 110M参数] --> B[量化 FP32→INT8] B --> C[模型大小减半 推理提速1.5倍] C --> D[非结构化剪枝] D --> E[剪掉30%参数 精度损失<1%] E --> F[知识蒸馏] F --> G[原始模型→学生模型] G --> H[压缩后模型 60M参数] style A fill:#ffebee style B fill:#ffcdd2 style C fill:#ef9a9a style D fill:#e57373 style E fill:#ef5350 style F fill:#f44336 style G fill:#e53935 style H fill:#d32f2f,color:#fff

折腾了一圈,最后采用了一个综合方案:

  • 量化:从 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

最终结果:

原始模型与压缩后模型在各指标上的对比,包括参数量、模型大小、推理延迟和分类准确率

指标原始模型压缩后模型变化
参数量110M60M-45%
模型大小420MB180MB-57%
推理延迟45ms28ms-38%
分类准确率92.3%90.8%-1.5%

虽然不是完美结果,但已经满足部署需求,而且精度损失在可接受范围内。

踩坑总结

这次折腾过程有几个关键的踩坑点:

  1. 不要迷信理论压缩比:理论上的压缩率往往假设硬件和软件完美支持,实际部署时各种限制会打折扣。

  2. 剪枝前要充分评估重要性:基于权重的简单评估容易误判,基于梯度或一阶导数的评估更可靠,但计算开销更大。

  3. 微调阶段要给足时间:剪枝破坏了模型结构,恢复需要更长时间的微调,不要急于求成。

  4. 硬件限制要提前确认:稀疏计算、量化加速这些特性,不是所有硬件都支持,提前确认能省很多时间。

  5. 综合方案往往优于单一方法:剪枝、量化、蒸馏这些手段,组合使用效果更好,但要注意各步骤的顺序和参数调节。

写在最后

模型剪枝不只是技术优化,更像在有限的资源约束下做取舍。有些参数删了可惜,但为了部署只能删;有些优化理论上很好,但工程落地困难只能放弃。

整个过程有点像装修房子:预算有限,空间有限,既要保留核心功能,又要控制在可接受的成本内。最后拿到的方案可能不是最优解,但一定是当前约束条件下的可行解。

这次折腾也让我意识到,边缘端部署和服务器端部署完全是两个世界。服务器端可以堆硬件、堆算力,但边缘端每个字节的内存、每毫秒的延迟都要精打细算。这种约束下的优化,反而更有意思一些。

如果下次再做类似的部署项目,我会更早地考虑硬件和软件的限制条件,把方案设计得更务实一些。毕竟,能跑起来的方案才是好方案。

版权声明: 本文首发于 指尖魔法屋-AI 剪枝踩坑记录https://blog.thinkmoon.cn/post/277-ai-pruning-structured-unstructured-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!