ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

基于模型压缩的识别算法Python源码:蒸馏与剪枝实战指南

基于模型压缩的识别算法Python源码:蒸馏与剪枝实战指南 简介这份资源是面向毕业设计与模型压缩入门者的Python代码仓库聚焦基于知识蒸馏与剪枝的识别算法实现适合具备一定深度学习基础、需要完成相关课题或实验对比的学生与开发者。压缩包共185个文件约4.03MB以79个py源码文件为核心辅以60个pyc编译文件、7个txt说明、2个json配置及多个训练日志与记录文件涵盖模型训练、剪枝配置、蒸馏记录等模块目录结构便于按实验阶段检索。资源围绕模型压缩技术展开包含不同数据集上的模型对比实验并涉及模型转换为Apple Silicon架构的实践可帮助读者理解蒸馏与剪枝的完整流程、复现对比结果并迁移到本地环境。目前已有113人学习适合作为毕业设计参考或模型压缩方向的练手项目。1. 识别模型压缩这件事为什么蒸馏加剪枝是性价比最高的起手式你训好了一个识别模型mAP 看着还行但一算参数量和推理延迟部署端直接摇头。这时候摆在面前的路无非几条换更小的骨干、量化、蒸馏、剪枝。换骨干意味着重新调参量化对某些算子支持不友好而蒸馏和剪枝可以叠加使用是大多数一线团队压缩识别算法的默认起手式。这个标题里的「基于模型压缩的识别算法python源码蒸馏和剪枝」本质就是一套把大模型能力迁移到小模型、再把小模型冗余结构砍掉的组合方案。它解决的核心问题是在精度掉点可控的前提下把模型体积和计算量压下来。适合谁适合已经有一个能跑通的识别模型分类、检测都算、想把它塞进边缘设备或移动端的工程师。蒸馏负责「教」剪枝负责「砍」两者顺序和参数没调好精度会崩得让你怀疑人生。下面按我实际落地的顺序把原理、代码和坑一条条拆开。2. 蒸馏先立住教师学生结构怎么搭、损失怎么配2.1 蒸馏为什么对识别任务有效知识蒸馏的核心逻辑是教师模型输出的软标签soft label里包含了类别间的相似性信息这些信息比硬标签更丰富。识别任务里比如一个动物分类模型教师对「狼」和「哈士奇」的输出分布是有重叠的这种暗知识dark knowledge能帮学生模型学到更平滑的决策边界。温度系数 T 就是用来放大这种软标签信息的T 越大分布越平滑暗知识暴露得越充分。但识别任务和纯分类有个区别检测任务里还有框回归分支蒸馏不能只蒸分类头。常见做法是对分类分支用 KL 散度蒸软标签对回归分支用 L2 蒸特征或直接蒸输出。我一般会先只蒸分类头看精度恢复情况再决定要不要加回归蒸馏。2.2 搭一个最小可跑的蒸馏训练脚本下面这段代码是一个分类识别任务的蒸馏训练骨架教师模型冻结学生模型正常训练损失由硬标签交叉熵和软标签 KL 散度加权组成。import torch import torch.nn as nn import torch.nn.functional as F class DistillLoss(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): # 硬标签损失学生直接对真实标签负责 hard_loss self.ce(student_logits, labels) # 软标签损失KL 散度注意两边都要除以温度 T soft_student F.log_softmax(student_logits / self.T, dim1) soft_teacher F.softmax(teacher_logits / self.T, dim1) soft_loss F.kl_div(soft_student, soft_teacher, reductionbatchmean) # 温度平方缩放保证梯度量级和硬损失可比 soft_loss soft_loss * (self.T ** 2) return self.alpha * soft_loss (1 - self.alpha) * hard_loss # 训练循环片段 teacher.eval() for images, labels in train_loader: with torch.no_grad(): teacher_logits teacher(images) student_logits student(images) loss distill_criterion(student_logits, teacher_logits, labels) loss.backward() optimizer.step() optimizer.zero_grad()逻辑说明temperature控制软标签平滑程度分类任务常用 3 到 5识别任务类别多可以适当调大。alpha是软损失权重0.7 是个稳妥起点学生容量越小越依赖软标签可以往 0.8 调。soft_loss * T**2这一步不能省否则温度一变损失量级就变了梯度会不稳定。参数说明教师模型必须切eval()并且用torch.no_grad()包住前向否则显存直接翻倍。学生和教师的输入预处理要完全一致归一化参数不同会让蒸馏效果大打折扣。2.3 教师学生容量差多少才合适这不是玄学是有边界的。教师比学生大 3 到 10 倍参数量是比较舒服的区间。差太小暗知识不够丰富蒸了等于没蒸差太大学生容量接不住精度反而比直接训还差。我踩过的坑是拿一个 100M 的教师去蒸一个 1M 的学生学生直接学崩最后还不如从头训。稳妥做法是先蒸一个中等大小的学生再拿这个学生当教师去蒸更小的这叫多级蒸馏。3. 剪枝再动手结构化与非结构化怎么选、怎么剪3.1 剪枝的两种路线和识别任务的适配剪枝分非结构化和结构化。非结构化剪枝是把权重矩阵里小的值置零稀疏度高但硬件不一定加速除非你用支持稀疏计算的推理引擎。结构化剪枝是直接砍通道、砍卷积核剪完模型结构真的变小通用硬件都能加速。识别算法部署到边缘设备我一般优先选结构化剪枝因为落地确定性高。剪枝粒度上通道剪枝channel pruning最常用。核心是给每个通道算一个重要性分数按分数排序砍掉最低的那批。重要性度量常见的有 L1 范数、BN 层缩放因子、泰勒展开。BN 缩放因子法Network Slimming实现简单且效果稳定推荐作为第一版方案。3.2 用 BN 缩放因子做通道剪枝的完整步骤第一步在训练时对 BN 层加 L1 正则让缩放因子稀疏化。第二步收集所有 BN 缩放因子按全局阈值或百分比确定剪枝比例。第三步真正砍通道并重建模型。下面是一个可复现的剪枝脚本核心部分。import torch import torch.nn as nn def collect_bn_scales(model): 收集所有 BN 层的缩放因子 scales [] for m in model.modules(): if isinstance(m, nn.BatchNorm2d): scales.append(m.weight.data.abs().clone()) return torch.cat(scales) def prune_model_by_threshold(model, threshold): 按阈值裁剪 BN 层通道返回裁剪掩码 masks {} for name, m in model.named_modules(): if isinstance(m, nn.BatchNorm2d): mask m.weight.data.abs() threshold masks[name] mask # 裁剪 BN 参数 m.weight.data m.weight.data[mask] m.bias.data m.bias.data[mask] m.running_mean m.running_mean[mask] m.running_var m.running_var[mask] m.num_features int(mask.sum()) return masks # 训练时加 L1 正则 def bn_l1_regularizer(model, lam1e-4): reg 0.0 for m in model.modules(): if isinstance(m, nn.BatchNorm2d): reg m.weight.abs().sum() return lam * reg逻辑说明collect_bn_scales把所有 BN 缩放因子拼成一个全局向量这样阈值是全局统一的避免逐层设阈值的麻烦。prune_model_by_threshold只处理了 BN 层本身实际剪枝还要同步裁剪前面的卷积输出通道和后面的卷积输入通道这部分需要根据模型结构手动或半自动处理是剪枝里最容易翻车的地方。参数说明lam是 L1 正则系数1e-4 到 1e-5 之间比较常见太大精度掉得厉害太小稀疏化不够。阈值 threshold 一般按全局缩放因子的百分位数来定比如砍掉 30% 通道就取 30% 分位数。3.3 剪枝后的微调策略剪完不微调精度基本没法看。微调学习率要用小学习率通常是初始学习率的十分之一甚至更低因为模型结构已经变了大学习率会把剩下的权重也带偏。微调轮数不用太多识别任务一般 10 到 20 个 epoch 就能恢复大部分精度。如果剪枝比例超过 50%微调可能也救不回来这时候要回退剪枝比例。4. 蒸馏和剪枝的先后顺序先蒸后剪还是边蒸边剪4.1 两种顺序的精度和耗时对比先蒸后剪先把大模型蒸到小模型再对小模型剪枝微调。优点是流程清晰每步可控缺点是剪枝后学生结构变了之前蒸的知识可能部分失效。边蒸边剪在蒸馏训练过程中逐步剪枝让教师持续指导学生适应剪后的结构。优点是精度通常更好缺点是实现复杂训练不稳定。我实测下来先蒸后剪在剪枝比例 30% 以内时精度损失很小工程上更好落地。边蒸边剪适合剪枝比例大、精度要求苛刻的场景但调参成本高。顺序精度保持实现难度适用剪枝比例先蒸后剪中等低30% 以内边蒸边剪较好高30% 到 60%只蒸不剪好低不剪只剪不蒸差低不推荐4.2 一个组合训练的配置模板如果你决定先蒸后剪训练分三段第一段纯蒸馏学生学教师第二段加 L1 正则继续蒸馏让 BN 缩放因子稀疏化第三段剪枝后微调此时教师可以继续用也可以不用。下面是一个配置模板。# 阶段一纯蒸馏 config_stage1 { epochs: 50, lr: 0.01, temperature: 4.0, alpha: 0.7, bn_l1_lambda: 0.0 } # 阶段二蒸馏 BN 稀疏化 config_stage2 { epochs: 30, lr: 0.005, temperature: 4.0, alpha: 0.7, bn_l1_lambda: 1e-4 } # 阶段三剪枝后微调 config_stage3 { epochs: 15, lr: 0.001, temperature: 4.0, alpha: 0.5, # 剪枝后可以降低软标签依赖 bn_l1_lambda: 0.0 }逻辑说明阶段二加 L1 正则后BN 缩放因子会向零靠拢为剪枝做准备。阶段三学习率降到 0.001因为结构变了需要精细调整。alpha在阶段三可以降到 0.5因为剪枝后学生结构已经固定硬标签的监督更重要。参数说明阶段二的bn_l1_lambda不要一开始就设大先跑几个 epoch 看缩放因子分布再调。阶段三如果精度恢复不理想可以把alpha调回 0.7 再试。5. 避坑与排查蒸馏剪枝里最容易翻车的五个点5.1 蒸馏损失不下降学生输出全是一类现象训练几个 epoch 后学生预测分布坍缩到某个类别KL 散度不降反升。原因温度 T 设得太小软标签不够平滑学生学不到暗知识或者alpha太大硬标签监督被淹没。解决先把 T 调到 5 以上alpha降到 0.5确认学生能正常学硬标签后再逐步加软损失权重。5.2 剪枝后模型直接报维度错误现象剪枝脚本跑完前向传播时报通道数不匹配。原因只剪了 BN 层没有同步裁剪相邻卷积层的输入输出通道。解决剪枝必须成组处理一个 BN 层对应前面卷积的输出通道和后面卷积的输入通道建议用现成的剪枝库或写一个依赖图来管理。5.3 微调后精度比剪枝前还低现象剪枝后微调了 20 个 epoch精度还是比剪枝前低很多。原因剪枝比例过大或者微调学习率太大把剩余权重带偏了。解决先降低剪枝比例从 10% 开始试微调学习率降到初始的十分之一并加 warmup。5.4 教师模型显存爆了现象蒸馏训练时显存占用是正常训练的两倍以上。原因教师模型前向没有用torch.no_grad()计算图被保留。解决教师前向必须包在with torch.no_grad():里并且教师模型切eval()模式关闭 dropout 和 BN 更新。5.5 BN 缩放因子稀疏化后剪枝阈值不会选现象L1 正则跑完缩放因子分布还是集中在某个区间没有明显的零附近聚集。原因L1 系数太小或者训练轮数不够。解决把bn_l1_lambda调大一个量级多跑 10 个 epoch观察缩放因子直方图等到零附近出现明显峰值再剪。6. 进阶技巧用敏感度分析决定每层剪多少一刀切按全局阈值剪枝对识别任务其实不够精细。不同层对剪枝的敏感度差别很大浅层卷积通常更敏感深层冗余更多。我后来习惯先做一轮敏感度分析逐层单独剪 10% 通道看精度掉多少掉得少的层多剪掉得多的层少剪甚至不剪。具体做法是写一个循环每次只剪一层跑一次验证集记录精度变化。这个过程比较耗时但能换来更精细的剪枝策略。下面是一个敏感度分析的骨架。def sensitivity_analysis(model, val_loader, prune_ratio0.1): 逐层剪枝记录精度变化 results {} base_acc evaluate(model, val_loader) for name, m in model.named_modules(): if isinstance(m, nn.BatchNorm2d): # 备份 backup m.weight.data.clone() # 剪掉该层最小的一部分通道 threshold torch.quantile(m.weight.data.abs(), prune_ratio) mask m.weight.data.abs() threshold m.weight.data[~mask] 0 # 模拟剪枝 acc evaluate(model, val_loader) results[name] base_acc - acc m.weight.data backup # 恢复 return results逻辑说明这里用置零模拟剪枝避免真正改结构跑完能恢复。prune_ratio是模拟剪枝比例一般取 0.1 到 0.2。results里记录的是每层剪枝后的精度下降量下降量小的层就是可以多剪的层。参数说明torch.quantile取的是分位数prune_ratio0.1表示砍掉该层最小的 10% 通道。实际剪枝时对敏感度低的层可以把比例提到 40% 甚至 50%敏感层保持在 10% 以内。最后说个我自己的习惯每次剪枝前一定先存一份完整模型权重剪枝脚本跑完先在小验证集上快速验证确认没有维度错误和精度崩塌再上全量微调。这个后悔药我吃过太多次亏才养成希望帮到你。本文还有配套的精品资源点击获取
返回列表