ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

模型压缩实战:蒸馏与剪枝源码解析及边缘部署优化

模型压缩实战:蒸馏与剪枝源码解析及边缘部署优化 简介这份资源是面向毕业设计与模型压缩入门者的Python代码仓库聚焦基于知识蒸馏与剪枝的识别算法实现适合具备一定深度学习基础、需要完成相关课题或复现压缩实验的学生与开发者。压缩包共185个文件约4.03MB以79个py源码文件为核心辅以60个pyc编译文件、若干txt说明、json配置、sh脚本及训练日志与记录文件覆盖模型训练、剪枝、蒸馏与结果记录等环节。内容涉及知识蒸馏、模型剪枝、不同数据集上的模型对比以及将模型转换为Apple Silicon架构等实践方向可帮助读者理解压缩流程、对照实验配置并复用训练脚本。目前已有113人学习适合作为毕设参考或模型压缩练手项目。1. 模型压缩识别算法源码包蒸馏加剪枝到底能压到什么程度一个训练集上 98% 准确率的识别模型部署到边缘设备上推理一次要 300 毫秒内存占用 200MB 以上这种场景做视觉识别的工程师基本都遇到过。模型压缩要解决的就是这个问题在不显著掉点的前提下把模型体积和推理耗时压下来。蒸馏和剪枝是两条最主流的路线前者让小模型学大模型的输出分布后者直接砍掉冗余的通道或权重。这个源码包把两条路线都做成了可跑的 Python 实现适合已经有一个能用的识别模型、想把它塞进更小硬件里的从业者。读完你应该能判断自己的场景该走哪条路、参数怎么设、哪里容易翻车。2. 蒸馏和剪枝的选型逻辑先搞清楚你的瓶颈在哪2.1 蒸馏适合什么场景剪枝适合什么场景蒸馏的本质是知识迁移。你有一个大模型教师它的 softmax 输出不只是「这个样本属于类别 A」而是「属于 A 的概率 0.85属于 B 的概率 0.12属于 C 的概率 0.03」。后面这些信息叫暗知识dark knowledge它编码了类别之间的相似性关系。小模型学生直接学硬标签只能学到「A 是对的」但学教师的软输出能学到「A 和 B 比较像和 C 差很远」。这就是为什么蒸馏出来的小模型往往比直接用硬标签训练的同结构模型高 1 到 3 个点。蒸馏适合的场景很明确你手头已经有一个精度不错但跑不动的大模型同时愿意接受一个结构更小的学生网络。学生网络的结构可以自己设计也可以直接用现成的轻量骨干比如 MobileNet 系列、ShuffleNet 系列。蒸馏不改变学生网络的结构它改变的是训练信号。剪枝的逻辑完全不同。它假设训练好的网络里有大量冗余参数把这些参数去掉精度不会掉太多。剪枝分两类非结构化剪枝是把单个权重置零产生稀疏矩阵理论上能压缩存储但实际推理加速需要硬件和推理引擎支持稀疏计算否则加速效果很有限结构化剪枝是直接砍掉整个卷积核或通道产生的是稠密的小模型通用硬件上就能加速。源码包里两种都实现了但如果你追求实际推理加速优先看结构化剪枝那条路。选型判断可以用一个简单规则如果你的瓶颈是「没有小模型可用但有大模型和训练数据」走蒸馏如果瓶颈是「已经有一个结构还行但参数冗余的模型」走剪枝如果两者都想要先剪枝再蒸馏或者交替做。2.2 源码包的整体结构和依赖环境拿到一个压缩源码包第一件事不是跑训练而是把目录结构和依赖关系摸清楚。常见的组织方式是按方法分目录每个方法下面有独立的模型定义、训练脚本和配置文件。依赖方面PyTorch 是主力框架版本建议 1.10 以上因为要用到torch.nn.utils.prune模块和一些较新的算子。其他依赖包括 numpy、tqdm、tensorboard可选看训练日志用。环境配置这一步很多人翻车。Python 版本建议 3.8 到 3.10太新的版本某些 PyTorch 轮子还没跟上。如果你用 conda创建一个独立环境是最稳妥的做法conda create -n model_compress python3.9 conda activate model_compress pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 pip install numpy tqdm tensorboard如果你用 venv 而不是 conda逻辑一样先隔离环境再装包。CUDA 版本根据你显卡驱动选cu118对应 CUDA 11.8老卡可能要用cu113或cu102。装完之后用一行命令验证import torch print(torch.__version__) print(torch.cuda.is_available())如果第二行输出False说明 CUDA 没配好后面训练会退回 CPU速度差几十倍。这是第一个要排查的点。2.3 蒸馏的核心实现损失函数和温度参数蒸馏的损失函数是两部分加权一部分是学生输出和教师软输出的 KL 散度另一部分是学生输出和真实标签的交叉熵。温度参数 T 控制软输出的平滑程度T 越大概率分布越平滑暗知识越丰富但太大会让分布接近均匀反而丢失信息。常见取值是 3 到 10 之间分类任务上 T4 或 T5 是比较稳的起点。下面是一个蒸馏损失的核心实现你可以直接对照源码包里的对应文件看import torch import torch.nn as nn import torch.nn.functional as F class DistillationLoss(nn.Module): def __init__(self, temperature4.0, alpha0.7): super().__init__() self.T temperature self.alpha alpha # 蒸馏损失权重 self.ce nn.CrossEntropyLoss() def forward(self, student_logits, teacher_logits, labels): # 软标签损失KL 散度注意要乘 T^2 保持梯度量级 soft_loss F.kl_div( F.log_softmax(student_logits / self.T, dim1), F.softmax(teacher_logits / self.T, dim1), reductionbatchmean ) * (self.T ** 2) # 硬标签损失 hard_loss self.ce(student_logits, labels) # 加权求和 return self.alpha * soft_loss (1 - self.alpha) * hard_loss这段代码里有两个容易忽略的点。第一soft_loss乘了T^2原因是 softmax 除以 T 之后梯度会缩小 T^2 倍不补回来蒸馏损失的梯度会太小训练不动。第二alpha控制软硬损失的比例经验值在 0.5 到 0.9 之间学生和教师差距越大alpha 可以适当调大让学生更多依赖教师的软信号。教师模型在蒸馏训练过程中要冻结参数并且切到 eval 模式。如果忘了切 evalBatchNorm 层会用当前 batch 的统计量导致教师输出不稳定学生学到的信号有噪声。这个坑很隐蔽因为 loss 不会报错只是最终精度差一截。2.4 剪枝的核心实现结构化通道剪枝的流程结构化剪枝的流程比蒸馏多几个步骤先训练一个基准模型然后评估每个通道的重要性按重要性排序去掉最不重要的通道最后对剪枝后的模型做微调恢复精度。重要性评估的准则有很多种L1 norm 是最简单也最常用的一个卷积核的权重绝对值之和越小说明这个核的输出越小越不重要。import torch.nn.utils.prune as prune def prune_conv_l1(model, amount0.3): 对模型中所有 Conv2d 层做 L1 结构化剪枝 for name, module in model.named_modules(): if isinstance(module, torch.nn.Conv2d): prune.ln_structured( module, nameweight, amountamount, n1, dim0 ) return model这里dim0表示按输出通道维度剪n1表示用 L1 norm。amount0.3表示剪掉 30% 的通道。注意 PyTorch 的prune模块默认是「重参数化」方式它给权重加了一个 mask并没有真正把通道删掉推理时计算量没变。要真正加速需要调用prune.remove把 mask 固化然后手动重建一个更窄的模型把保留的通道权重拷过去。源码包里通常会有compress_model或rebuild_model这样的函数来做这件事你重点看那部分。剪枝率不能一次设太高。经验做法是迭代剪枝每次剪 10% 到 20%微调几个 epoch再剪下一轮。一次性剪 50% 以上精度基本会崩微调也救不回来。这个血泪经验在结构化剪枝上尤其明显。3. 从零跑通蒸馏训练数据准备、教师加载和学生配置3.1 数据加载和教师模型加载的注意事项数据部分用标准的torchvision.datasets.ImageFolder或自定义 Dataset 都行关键是训练集和验证集的划分要固定随机种子否则每次跑出来的精度波动会让你误以为是蒸馏参数的问题。教师模型的加载有两种方式如果教师是 torchvision 自带的预训练模型直接models.resnet50(pretrainedTrue)就行如果是自己训练的 checkpoint用torch.load加载 state_dict注意map_location要设对否则 GPU 上存的模型在 CPU 环境加载会报错。import torch import torchvision.models as models # 加载教师模型 teacher models.resnet50(pretrainedFalse) teacher.load_state_dict(torch.load(teacher_best.pth, map_locationcpu)) teacher.eval() for param in teacher.parameters(): param.requires_grad False # 学生模型用更轻的骨干 student models.resnet18(pretrainedFalse)教师模型一定要先eval()再冻结梯度。顺序反了不影响功能但养成先 eval 的习惯能避免很多玄学问题。学生模型的结构选择上分类任务用 ResNet18 或 MobileNetV2 都行检测任务学生骨干要和检测头匹配不能随便换。3.2 训练循环里蒸馏 loss 的接入方式训练循环和普通训练几乎一样唯一区别是每个 batch 要同时跑教师和学生然后把两个 logits 一起送进蒸馏损失函数。下面是一个最小训练循环的骨架optimizer torch.optim.SGD(student.parameters(), lr0.01, momentum0.9) criterion DistillationLoss(temperature4.0, alpha0.7) for epoch in range(num_epochs): student.train() for images, labels in train_loader: images, labels images.cuda(), labels.cuda() with torch.no_grad(): teacher_logits teacher(images) student_logits student(images) loss criterion(student_logits, teacher_logits, labels) optimizer.zero_grad() loss.backward() optimizer.step()教师前向要包在torch.no_grad()里否则会占额外显存batch size 大一点就 OOM。学习率方面蒸馏训练的学习率可以比从头训练稍小因为学生有教师引导不需要太大的步长去探索。常见设置是 0.01 到 0.05 之间配合 cosine 衰减。3.3 蒸馏效果验证看哪些指标怎么对比验证不能只看学生模型的 top-1 准确率。你至少要对比三组数字学生从头训练无蒸馏的精度、学生蒸馏后的精度、教师模型的精度。如果蒸馏后的学生比无蒸馏学生高不到 0.5 个点说明蒸馏没起作用要检查温度、alpha 和教师输出是否正常。一个快速检查方法是打印教师 softmax 输出的熵如果熵接近 0分布太尖锐说明温度太低暗知识没提取出来如果熵接近 log(类别数)分布太均匀说明温度太高。另外蒸馏训练收敛通常比普通训练慢因为软标签的梯度信号更平滑。不要跑几个 epoch 看 loss 没降就放弃至少跑完完整的学习率衰减周期再判断。4. 剪枝实操重要性评估、迭代剪枝和微调恢复4.1 用 L1 norm 做通道重要性排序在动手剪之前先把每一层卷积核的 L1 norm 算出来看看分布。如果某一层的 norm 分布很集中说明这层冗余度高可以多剪如果分布很分散说明每个通道都在起作用剪多了会掉点。def compute_channel_importance(model): importance {} for name, module in model.named_modules(): if isinstance(module, torch.nn.Conv2d): # 每个输出通道的 L1 norm norm module.weight.data.abs().sum(dim(1, 2, 3)) importance[name] norm.cpu().numpy() return importance拿到 importance 之后可以画个直方图或者按层打印 min/max/mean。经验上靠近输入的层和靠近输出的层对剪枝更敏感中间层冗余度更高。所以剪枝率可以分层设置浅层剪 10% 到 20%中间层剪 30% 到 40%深层剪 20% 到 30%。4.2 迭代剪枝的完整流程和微调策略迭代剪枝的伪代码逻辑是这样的训练基准模型 → 评估重要性 → 剪一层或几层 → 微调 → 重复直到达到目标压缩率。微调的学习率要比初始训练小一个量级通常用 0.001 左右跑 10 到 20 个 epoch。微调数据用全部训练集不要只用子集否则恢复不充分。def iterative_prune(model, train_loader, target_sparsity0.5, step0.1): current_sparsity 0.0 while current_sparsity target_sparsity: # 剪一步 model prune_conv_l1(model, amountstep) # 微调 fine_tune(model, train_loader, epochs10, lr0.001) current_sparsity step print(fSparsity: {current_sparsity:.1f}, Acc: {evaluate(model):.4f}) return model每次剪完必须微调不能连续剪多次再一起微调那样精度掉下去就回不来了。微调的时候建议用比初始训练更小的学习率加上 warmup让模型慢慢适应结构变化。4.3 剪枝后模型重建从 mask 到真正的窄模型前面提过PyTorch 的prune模块只是加 mask不改变实际计算图。要真正加速必须重建模型。重建的逻辑是对每一层找出被保留的通道索引创建一个输出通道数等于保留数量的新 Conv2d把对应权重拷过去。如果下一层的输入通道依赖上一层的输出通道下一层的输入也要同步裁剪。这就是为什么结构化剪枝比非结构化剪枝麻烦它涉及层与层之间的依赖关系。def rebuild_conv(conv, keep_indices): 根据保留索引重建卷积层 out_channels len(keep_indices) in_channels conv.in_channels new_conv torch.nn.Conv2d( in_channels, out_channels, kernel_sizeconv.kernel_size, strideconv.stride, paddingconv.padding, biasconv.bias is not None ) new_conv.weight.data conv.weight.data[keep_indices].clone() if conv.bias is not None: new_conv.bias.data conv.bias.data[keep_indices].clone() return new_conv重建之后要跑一遍验证集确认精度和剪枝前一致或只差零点几个点然后再做后续微调。如果重建后精度暴跌大概率是通道索引对错了或者 BatchNorm 的 running_mean/running_var 没有同步裁剪。BatchNorm 的裁剪经常被忽略但它和卷积输出通道是一一对应的卷积剪了哪些通道BN 就要剪哪些。5. 避坑与排查蒸馏剪枝里最容易翻车的五个地方5.1 蒸馏 loss 不下降学生精度和从头训练一样现象训练日志里 loss 在降但验证集精度和不用蒸馏的学生模型几乎一样甚至更低。原因最常见的是教师模型没有切 eval 模式BatchNorm 统计量在训练中不断变化教师输出不稳定。其次是温度参数设得太小比如 T1软标签退化成硬标签暗知识完全丢失。解决确认教师模型在蒸馏前调用了eval()并且冻结了梯度。温度从 T4 开始试观察教师输出的熵是否在合理范围。如果还不行检查 alpha 是不是设得太小软损失被硬损失淹没了。5.2 剪枝后模型推理速度没变现象剪了 40% 的通道验证精度也还行但推理耗时和剪枝前一样。原因用了 PyTorch 的prune模块但没有调用prune.remove和重建模型实际计算图没变只是权重被 mask 了。另一种可能是用了非结构化剪枝产生稀疏矩阵但推理引擎不支持稀疏加速。解决结构化剪枝后必须重建模型把 mask 变成真实的窄层。非结构化剪枝如果目标平台不支持稀疏计算就不要指望推理加速它只省存储。5.3 微调后精度恢复不到剪枝前水平现象剪枝率 30%微调 20 个 epoch精度还是比基准低 3 个点以上。原因剪枝率一次设太高或者微调学习率太大导致模型在窄结构上震荡。也可能是剪枝时把关键层剪太狠了比如第一层或最后一层。解决降低单次剪枝率改成每次 10% 迭代剪。微调学习率降到初始训练的十分之一加 warmup。检查各层剪枝率是否均匀浅层和深层少剪。5.4 蒸馏训练显存不够batch size 上不去现象教师和学生同时前向显存占用是单模型的两倍多batch size 只能设很小训练不稳定。原因教师前向没有包在torch.no_grad()里PyTorch 为教师也建了计算图显存直接翻倍。解决教师前向必须用with torch.no_grad():包住。如果显存还是紧张可以把教师输出提前算好存成文件训练时直接读 logits这样训练阶段只需要跑学生。5.5 剪枝后 BatchNorm 统计量不匹配导致精度崩现象重建模型后精度从 95% 掉到 60% 多但权重明明是对应拷贝的。原因卷积通道剪了但后面的 BatchNorm 没有同步裁剪通道数对不上或者 running_mean/running_var 还是旧的。解决重建卷积的同时重建对应的 BatchNorm把保留通道的 running_mean、running_var、weight、bias 一起拷过去。重建后在训练集上跑几百个 batch 让 BN 重新统计再做微调。6. 进阶技巧蒸馏和剪枝交替做以及一个验证压缩收益的硬指标单独做蒸馏或剪枝压缩率到 2 到 3 倍之后就会遇到瓶颈。继续剪精度崩继续蒸学生容量不够学不动。这时候可以交替做先剪枝得到一个中间模型用它当教师去蒸馏一个更小的学生学生训练完再剪一轮。这个流程能把压缩率推到 5 倍以上但每一步都要验证精度不能跳步。一个具体的交替流程是这样的基准模型剪枝 30% 得到模型 A微调恢复精度用模型 A 当教师蒸馏一个通道数减半的学生模型 B模型 B 再剪枝 20%微调。最终模型 B 的体积大概是基准的 25% 到 30%精度通常能保持在基准的 95% 以上。这个数字不是绝对的取决于你的任务难度和基准模型冗余度。验证压缩收益不能只看参数量。参数量少不代表推理快因为推理速度还受内存访问、算子实现、硬件并行度影响。真正要看的指标是在目标硬件上的单次推理延迟毫秒、峰值内存占用MB、以及精度。这三个数字放在一起才能判断压缩方案值不值得上。我一般会做一个表格把基准模型、蒸馏模型、剪枝模型、交替压缩模型的这三项指标列出来一目了然。模型版本参数量(M)推理延迟(ms)峰值内存(MB)精度(%)基准模型25.64521096.2蒸馏学生11.22210594.8剪枝模型15.32813095.1交替压缩7.8167893.5这张表里的数字是示意你跑自己的任务时把真实数字填进去。重点看推理延迟和精度的比值如果延迟降了一半但精度只掉 1 个点这个方案就值得上如果延迟只降 20% 但精度掉 3 个点就要重新考虑剪枝策略或者换更轻的骨干。还有一个容易忽略的点压缩后的模型在不同硬件上的表现可能完全不同。在服务器 GPU 上剪枝带来的加速可能不明显因为 GPU 并行度高小模型的算力利用率本来就低但在边缘设备上剪枝的加速效果会明显得多。所以验证一定要在目标硬件上做不能拿服务器 GPU 的数字去推断边缘设备的表现。我自己做压缩项目时养成了一个习惯每做完一轮压缩先把模型导出成 ONNX 或者 TorchScript在目标推理引擎上跑一遍 benchmark再决定要不要继续压。因为训练框架里的推理速度和部署后的推理速度经常对不上提前暴露问题比部署后再返工成本低得多。希望帮到你。本文还有配套的精品资源点击获取
返回列表