把卷积网络或Transformer直接部署到树莓派、Jetson Nano或手机NPU上时,往往会遇到同一个问题:模型在服务器上推理一次只有几毫秒,搬到边缘端后单帧耗时超过几百毫秒,甚至因内存不足直接加载失败。算力、功耗和散热三重限制下,单纯优化推理框架或降低批次大小很难获得质变,更有效的做法是从模型本身移除冗余。模型剪枝和知识蒸馏是两条主流压缩路径:前者删除对输出影响较小的权重或通道,后者让小模型学习大模型的软输出分布。将两者组合使用,通常能得到比单独使用更小的体积和更好的精度。

一、模型剪枝:从权重稀疏到通道裁剪
深度神经网络普遍存在参数冗余,很多连接对最终输出的贡献极低。模型剪枝的目标就是识别并移除这些低价值参数。它大致分成两类:非结构化剪枝以单个权重为单位,把绝对值较小的权重置零,得到的稀疏矩阵理论上能减少存储和乘法次数;结构化剪枝则以卷积核、通道或整个层为单位,直接生成更小的密集网络。前者的稀疏度可以很高,但通用处理器和GPU对稀疏矩阵的加速并不稳定,边缘端推理框架的支持也有限。后者虽然剪枝粒度更粗,但裁剪后模型仍然是规则的密集结构,更容易被ONNX Runtime、TensorRT、Core ML等后端高效执行。
判断哪些通道该剪,常用指标包括权重绝对值、梯度信息、BatchNorm层的gamma缩放系数以及Taylor展开带来的损失变化。其中BatchNorm gamma是最易落地的方案:每个通道都有一个可学习的gamma,它接近零说明该通道输出幅度很小,对下一层贡献弱。剪枝流程通常是先训练一个完整模型,对目标层的gamma排序,按比例移除低分通道,然后对剪枝后的模型重新微调若干轮。为了避免一次性损失过大,也可以把剪枝比例拆成多次迭代执行。
下面是一个简化的通道剪枝示例,基于BatchNorm gamma计算保留通道索引,并同步裁剪前后卷积层权重:
import torch
def select_keep_indices(bn_module, prune_ratio):
gamma = bn_module.weight.data.abs()
num_channels = gamma.shape[0]
num_keep = max(1, int(num_channels * (1.0 - prune_ratio)))
keep_indices = torch.topk(gamma, num_keep).indices
return sorted(keep_indices.tolist())
def apply_channel_prune(conv_prev, bn_module, conv_next, keep_indices):
conv_prev.weight.data = conv_prev.weight.data[keep_indices, :, :, :]
conv_prev.out_channels = len(keep_indices)
bn_module.weight.data = bn_module.weight.data[keep_indices]
bn_module.bias.data = bn_module.bias.data[keep_indices]
bn_module.running_mean.data = bn_module.running_mean.data[keep_indices]
bn_module.running_var.data = bn_module.running_var.data[keep_indices]
bn_module.num_features = len(keep_indices)
conv_next.weight.data = conv_next.weight.data[:, keep_indices, :, :]
conv_next.in_channels = len(keep_indices)
这段代码展示的是最基础的操作,真实模型中还要处理残差连接的通道对齐、分组卷积和深度可分离卷积的约束。例如ResNet的残差分支要求输出通道数不能随意变化,剪枝时需要跨层协调;MobileNet的深度可分离卷积中,逐通道卷积的输出通道与逐点卷积的输入通道必须严格一致。忽略这些结构约束,重建出来的模型前向传播会直接报维度错误。
剪枝比例需要根据层敏感度设置,不能全局统一。浅层通常提取边缘和纹理等基础特征,对通道裁剪更敏感;深层特征冗余较高,可设置更大的剪枝比例。实践中可以先做一次逐层敏感性分析,记录每一层在不同剪枝比例下的精度变化,再分配各层压缩预算。这样比盲目全局剪枝稳定得多。
二、知识蒸馏:软标签里的隐藏信息
一个训练好的大模型输出不仅是最大概率类别,还包含丰富的类间关系。以手写数字识别为例,数字3的图片经过softmax后,可能除了3之外还会给8和5分配一点概率,因为它们的形态有相似之处。硬标签只告诉学生模型这张图是3,而软标签还告诉它8比0更像3。这些额外信息能帮助学生模型用更少参数捕捉类别边界,缩小与大模型之间的泛化差距。
软标签由温度参数T控制。原始softmax在T等于1时分布较为尖锐,大模型对错误类别的概率可能趋近于零。增大T会让分布变平滑,让次要类别获得更多权重,但也可能把噪声放大。典型T取值范围是2到10,需要根据任务复杂度和教师模型置信度调节。温度越高,学生模型需要更多训练轮次来吸收平滑后的信息,因此不建议一开始就设得过大。
蒸馏训练的损失函数通常包含两项:一项是学生软输出和教师软输出之间的KL散度,另一项是学生输出与真实硬标签之间的交叉熵。KL散度乘以T的平方是为了抵消温度缩放对梯度幅值的影响。两者通过超参数alpha加权,alpha越大越依赖教师软标签,越小越偏向数据集原始标签。下面是PyTorch实现:
import torch
import torch.nn.functional as F
def distillation_loss(student_logits, teacher_logits, labels, T=4.0, alpha=0.7):
soft_student = F.log_softmax(student_logits / T, dim=1)
soft_teacher = F.softmax(teacher_logits / T, dim=1)
distill = F.kl_div(soft_student, soft_teacher, reduction='batchmean') * (T * T)
hard = F.cross_entropy(student_logits, labels)
return alpha * distill + (1.0 - alpha) * hard
def train_step(student, teacher, optimizer, images, labels):
student.train()
teacher.eval()
optimizer.zero_grad()
student_logits = student(images)
with torch.no_grad():
teacher_logits = teacher(images)
loss = distillation_loss(student_logits, teacher_logits, labels)
loss.backward()
optimizer.step()
return loss.item()
教师模型必须设置为eval模式并关闭梯度更新,否则会浪费大量显存和计算。学生模型的容量不能过小,否则即使软标签质量很高,也没有足够参数去拟合这些分布。选择学生结构时,可以参考教师模型的层数和通道数进行等比缩小,例如ResNet-34教师对应ResNet-18学生,或者MobileNetV3-Large对应MobileNetV3-Small。
三、剪枝与蒸馏如何协同:先压缩再学习
剪枝和蒸馏并不是竞争关系。剪枝负责减少参数和计算量,但剪掉通道后会损失一部分表示能力;蒸馏则用教师模型提供的软标签帮助剪枝模型重新学习被削弱的知识。二者配合的常见顺序是先剪枝后蒸馏:先训练教师模型,再对一个较大的学生结构做结构化剪枝,得到紧凑目标结构,最后用蒸馏损失微调这个紧凑模型。这样做的好处是训练过程中模型结构已经固定,不会因为后续剪枝再次引入精度波动。
另一种顺序是先蒸馏到小模型,再对小模型剪枝。这种方式更轻量,但剪枝后可能需要额外微调,整体精度恢复不如前者稳定。如果教师模型本身很大,还可以把中间层特征也加入蒸馏目标,让学生模型同时模仿浅层和深层表示。工程上更推荐从基础方案开始:教师模型用标准交叉熵训练,学生模型选择目标算力允许的结构,剪枝后再进行软标签蒸馏。
下面给出组合训练的大致流程,假设已经完成教师训练和学生通道选择:
# 1. 训练教师模型
teacher = train_big_model()
# 2. 构建学生模型并结构化剪枝
student = build_small_model()
keep = select_keep_indices(student.bn1, prune_ratio=0.3)
apply_channel_prune(student.conv1, student.bn1, student.conv2, keep)
# 3. 使用蒸馏损失微调剪枝后的学生
for epoch in range(finetune_epochs):
for images, labels in train_loader:
train_step(student, teacher, optimizer, images, labels)
# 4. 导出模型供边缘端推理
torch.onnx.export(student, dummy_input, 'pruned_student.onnx')
实际工程中应根据模型结构重新构造一个新的小网络,而不是在原始模块上原地修改参数。原地改完权重后,优化器状态、学习率调度和结构定义可能不同步,后续训练容易出错。建议把保留通道索引保存下来,写一个重建函数生成新的卷积层和BatchNorm层,再加载裁剪后的权重。这样后续微调、量化和导出都会更干净。
训练完成后还要检查剪枝后的模型在目标推理后端上是否真的变快。某些深度学习框架对动态形状输入支持有限,如果剪枝导致不同层输入通道不规律,导出ONNX后可能出现算子形状推断失败。保持网络结构对称、通道数为8或16的整数倍,通常更容易通过推理引擎的优化。
四、部署评估:别只看参数量和FLOPs
压缩模型时很容易过度关注参数量和FLOPs,但这两个指标不能完全代表边缘设备上的推理速度。FLOPs只反映理论乘加次数,实际耗时还受内存带宽、缓存命中率、算子调度和并行度影响。例如非结构化剪枝可以把FLOPs降得很低,但如果推理后端没有针对稀疏矩阵的高效实现,实际延迟甚至可能比原始密集模型还高。结构化剪枝虽然FLOPs下降幅度相对保守,但更容易获得稳定的硬件加速。
边缘设备部署还需要关注内存占用峰值和算子支持情况。模型加载后要同时容纳权重、中间特征图和运行时缓冲区,手机NPU或单片机内存通常比服务器小几个数量级。部署前应在目标设备上跑真实输入,记录首次推理和后续推理的延迟、内存峰值以及精度。下表给出一个示意性对比,实际数值随数据集和硬件变化:
| 模型版本 | 参数量 | FLOPs | CPU延迟 | 精度 |
|---|---|---|---|---|
| 原始ResNet-18 | 11.69M | 1.82G | 32ms | 71.2% |
| 剪枝后 | 4.31M | 0.71G | 14ms | 70.1% |
| 蒸馏后 | 4.31M | 0.71G | 14ms | 70.9% |
| 剪枝加蒸馏 | 4.31M | 0.71G | 14ms | 71.0% |
从表格可以看出,剪枝主要降低参数和计算量,但精度会出现一定回落;蒸馏在不增加推理成本的情况下把精度拉回接近原模型。两者结合后,模型在边缘端延迟下降一半以上,精度损失控制在一个百分点以内。如果还需要进一步压缩,可以在蒸馏微调后接入INT8量化,将权重和激活从32位浮点转为8位整数,通常能在支持INT8算子的NPU上再获得一到两倍加速。
部署前建议按以下顺序做工程验证:
- 优先结构化剪枝,确保推理后端能识别规则密集结构。
- 蒸馏温度从4开始调,再在2到10之间搜索,避免一次性设置过大。
- 同时记录精度、延迟和内存峰值,三者缺一不可。
- 在目标设备上跑真实推理,用桌面GPU延迟推算边缘性能往往偏差很大。
模型压缩没有一套参数适合所有任务。剪枝比例、蒸馏温度、学生结构和微调轮数都需要根据实际硬件约束和目标精度反复调整。把剪枝的结构精简能力和蒸馏的知识迁移能力结合起来,通常能在有限算力下获得更实用的边缘端模型。