ARTICLE DETAIL

资讯详情

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

VGGNet从138M到3M:稀疏训练与通道剪枝实战

VGGNet从138M到3M:稀疏训练与通道剪枝实战 简介面向模型压缩与剪枝实战的VGGNet完整资料包适合希望掌握结构化剪枝流程的深度学习者与部署工程师。资源围绕VGGNet的稀疏训练、剪枝、微调全流程展开涵盖训练脚本、稀疏因子引入、BN层权重排序与阈值选取、逐层掩码生成、新网络结构重建及微调实现最终目标是将模型压缩至约3M。包内共2000个文件以1992张过程可视化图片为主其中包含训练曲线、BN层权重直方图、各层掩码图与剪枝前后结构对比便于对照每个阶段的权重分布和剪枝效果另有5个Python脚本负责训练与剪枝逻辑2个JSON文件保存配置与结果1个TXT文件提供使用说明。压缩包整体约871MB所有文件按训练、稀疏化、剪枝、微调等模块分层存放目录结构清晰可直接复现实验。目前已有957人学习下载对正在研究模型压缩或需要精简图像分类模型的读者是兼具原理演示与代码参考的实操型资源。1. 3M的VGGNet怎么来的先算清138M参数到3M的账VGG16的官方权重有约138M个参数用float32存下来是553MB而标题里的“3M”通常指的是模型文件小于3MB——这意味着要压缩将近180倍听起来像压缩率比赛里才敢定的目标。但如果你把VGGNet的全连接层拆掉、用L1稀疏把BN的gamma训练到接近0、再按通道裁剪后微调这个数字在CIFAR-10这类单卡就能跑的数据集上是可以够到的。这套流程不是只针对VGGNet它背后的“稀疏训练-剪枝-微调”三段式是可以原样搬到ResNet和MobileNet上的通用方法论。适合两类人一类是手上只有单卡或CPU、要做端侧部署的工程师另一类是刚接触模型压缩、想把结构化剪枝的每个环节真正跑通而不是只看论文的新手。下面所有操作都基于PyTorch原生API不依赖任何魔改框架。2. 训练一个可剪的VGGNet换上GAP加L1稀疏基线决定裁剪上限2.1 VGGNet为什么先改结构再训练全连接层吃掉90%参数VGG16的参数量大头不在卷积而在最后三个全连接层。第一个FC的输入是7×7×51225088个神经元输出4096这一层的权重就有25088×4096约102M参数再加上后面两个FC全连接部分合计约123M占整个模型的89%。如果带着这三个FC层去做通道剪枝你会发现卷积通道都剪完了参数量还是几十M因为FC层根本不参与通道维度的裁剪。所以剪VGGNet的第一步不是调参而是改结构把三层FC换成全局平均池化GAP单层线性分类头。这个替换在CIFAR-10这类小分辨率数据集上损失很小。GAP把特征图直接压成1×1等价于让最后卷积层输出的每个通道都参与分类决策既消掉了最占参数的FC层又天然带了正则效果。替换后conv部分约有14.7M参数这才是通道剪枝能真正发挥作用的战场。我一般会在训练前就把网络定义成这种结构而不是先训完整版再改因为剪枝后的微调是从剪完的state_dict继续原版FC的权重反正用不上留着只会干扰你对“3M目标”的进度判断。import torch import torch.nn as nn class VGGSlim(nn.Module): def __init__(self, cfg, num_classes10): super().__init__() layers [] in_ch 3 for v in cfg: if v M: layers.append(nn.MaxPool2d(2)) else: layers [ nn.Conv2d(in_ch, v, 3, padding1), nn.BatchNorm2d(v), nn.ReLU(inplaceTrue) ] in_ch v self.features nn.Sequential(*layers) self.gap nn.AdaptiveAvgPool2d(1) self.classifier nn.Linear(cfg[-1], num_classes) def forward(self, x): x self.features(x) x self.gap(x).flatten(1) return self.classifier(x) # 略深一点的结构CIFAR-10上约88%左右基线 cfg [64, M, 128, M, 256, 256, M, 512, 512, M, 512, 512, M] model VGGSlim(cfg, num_classes10)这段结构里所有Conv2d后面都紧跟BatchNorm2d这是后面做通道剪枝的前提。BatchNorm的gamma即weight参数会被训练成通道重要性的指示器如果某个通道的gamma趋近于0说明这个通道的输出经过归一化后几乎不贡献信息可以剪掉。这里的cfg沿用了VGG16的卷积深度但去掉了三层FC。注意MaxPool后接的层数没有变化所以特征图的感受野逻辑和原版VGG基本一致CIFAR-10这种32×32输入下不会因为GAP引入严重的精度损失。2.2 稀疏训练怎么写给BN的gamma加L1正则通道剪枝的前提是让网络自己长出“可剪性”这就是稀疏训练要做的事。常规训练只优化交叉熵BN的gamma会分布在一个较宽的范围内剪枝时很难定一个阈值。做法是在原有loss上叠加一个L1惩罚项惩罚目标是所有BN层gamma的绝对值之和。L1范数的特性是会把一部分gamma精确推向0训练结束后通道的重要性分布会变成两极分化——一部分通道gamma接近0另一部分保持较大值阈值就好选了。实现的改动非常小不需要改网络结构只需要在训练循环里多算一项。但要留意这个L1项的系数它直接决定稀疏化的强度。系数太小gamma拉不动训练完和没加正则一样系数太大所有通道都被压扁精度先崩了。我一般从1e-4起步观察gamma的分布再微调。import torch.nn.functional as F def train_one_epoch(model, loader, optimizer, args): model.train() total_loss 0.0 for images, targets in loader: images, targets images.to(args.device), targets.to(args.device) optimizer.zero_grad() logits model(images) ce F.cross_entropy(logits, targets) # L1稀疏项对所有BN的gamma求绝对值之和 l1_reg 0.0 for m in model.modules(): if isinstance(m, nn.BatchNorm2d): l1_reg torch.abs(m.weight).sum() loss ce args.sparsity * l1_reg loss.backward() optimizer.step() total_loss loss.item() return total_loss / len(loader)weight_decay不要同时加到BN的gamma上否则L1项会和权重衰减互相拉扯gamma稀疏化的速度会变得很难判断。常见做法是weight_decay只作用于卷积权重和全连接权重BN的weight和bias在优化器里单独设weight_decay0。SGD的momentum保持0.9即可。如果你用AdamL1稀疏项对学习率的敏感度更高sparsity系数要再往下降这也是有人反馈“加了稀疏训练精度掉得厉害”的常见原因。2.3 基线训练超参batch size、学习率与稀疏系数怎么配CIFAR-10上这套VGGSlim的参考配置是batch size 128SGD初始学习率0.1momentum 0.9weight_decay 5e-4训练120个epoch在第40/60/80个epoch把学习率除以10。sparsity系数从第5个epoch开始生效前5个epoch只训交叉熵让网络先warm up。这个warm up很重要——如果从第1个epoch就加L1惩罚梯度方向会被“压小gamma”和“拟合数据”两个目标同时拽着前期收敛很慢而且后面不容易恢复。训练结束后先别急着剪先打印gamma的分布看稀疏化效果。把模型里所有BN的gamma取值收集起来计算接近0的比例。如果训练正确你应该能看到明显的双峰分布一批gamma集中在0附近一批在0.5到1以上。用NumPy做个简单统计import numpy as np gammas [] for m in model.modules(): if isinstance(m, nn.BatchNorm2d): gammas.append(m.weight.data.abs().cpu().numpy()) gammas np.concatenate(gammas) print(fgamma范围: [{gammas.min():.4f}, {gammas.max():.4f}]) print(f小于0.01的比例: {(gammas 0.01).mean():.2%})如果“小于0.01的比例”只有个位数百分比说明sparsity系数太小可以把1e-4提到3e-4或5e-4重训如果比例超过60%说明压太狠基线精度已经受伤往回调系数。这一步是剪枝前的体检体检不过关就不要往下走。基线精度决定了剪枝上限——我在CIFAR-10上把基线训到约88%剪掉70%通道后微调能回到86%左右如果基线只有82%剪完大概率掉到80%以下。3. 从稀疏模型到瘦身模型通道剪枝的完整流程与阈值选择3.1 剪枝阈值怎么定全局阈值与按层比例二选一稀疏训练完成后模型里每个BN的gamma代表对应通道的重要性。剪枝就是把gamma低于某个阈值的通道整条删掉对应的卷积输出通道删掉下一层卷积的输入通道也删掉同时BN层的参数、running_mean、running_var都要按索引同步删。这个操作叫结构化剪枝剪完的模型是“矮胖”的——通道数变少但网络结构是完整的可以直接用普通推理框架加载。阈值选取有两种常用策略。全局阈值是把所有BN的gamma合并成一个大数组取某个分位数作为阈值比如剪掉70%就是取第30百分位数。这种做法简单但浅层通道往往比深层通道更敏感全局阈值可能把浅层剪过头。按层比例则是每层独立指定保留比例浅层多留、深层多剪实际操作更稳但需要自己调每层的比例。我一般先从全局阈值开始目标剪枝率0.7跑完看每层被剪的比例如果发现某些层被清掉了超过90%就改用按层比例把浅层保留率上调。第一层卷积输入RGB的3通道我通常完全不剪因为3×3×3的输入通道本身只有3个剪掉后对特征提取影响太大。3.2 通道裁剪的代码实现重建Conv2d与BN层通道剪枝最容易出错的地方是索引对齐。剪掉输出通道时当前卷积的weight按输出维删紧随其后的BN参数也按同一份mask删而下一层卷积的输入通道由当前层的输出通道决定所以还要把下一层卷积weight的输入维也按同一份mask删。这意味着每一层的输出mask会作为下一层的输入mask传递下去写代码时必须把这条链锁对准。下面的函数遍历features里的卷积和BN逐层重建。核心思路是碰到卷积先存下它的引用碰到BN时利用当前累积的输入mask和这一层算出的输出mask重建一个通道数收缩后的Conv2d和BatchNorm2d。def prune_vgg_channel(model, thresh): old_layers list(model.features) new_layers [] in_keep torch.ones(3, dtypetorch.bool) # 输入RGB通道不剪 pending_conv None for m in old_layers: if isinstance(m, nn.Conv2d): pending_conv m elif isinstance(m, nn.BatchNorm2d): out_keep m.weight.data.abs() thresh old_w pending_conv.weight.data # 输出维用out_keep输入维用in_keep new_conv nn.Conv2d( int(in_keep.sum()), int(out_keep.sum()), pending_conv.kernel_size, pending_conv.stride, pending_conv.padding ) new_conv.weight.data.copy_( old_w[out_keep][:, in_keep, :, :] ) new_bn nn.BatchNorm2d(int(out_keep.sum())) new_bn.weight.data.copy_(m.weight.data[out_keep]) new_bn.bias.data.copy_(m.bias.data[out_keep]) new_bn.running_mean.copy_(m.running_mean[out_keep]) new_bn.running_var.copy_(m.running_var[out_keep]) new_layers [new_conv, new_bn] in_keep out_keep # 当前输出成为下一层的输入mask elif isinstance(m, nn.ReLU): new_layers.append(nn.ReLU(inplaceTrue)) elif isinstance(m, nn.MaxPool2d): new_layers.append(m) model.features nn.Sequential(*new_layers) return model # 按百分位数计算全局阈值剪掉70%通道 all_gamma np.concatenate(gammas) thresh np.percentile(all_gamma, 70) pruned_model prune_vgg_channel(model, thresh)这段代码有两个关键点。一是new_conv.weight.data.copy_的索引顺序先取out_keep再取in_keep因为PyTorch的Conv2d权重shape是[out_channels, in_channels, kh, kw]两个mask作用在不同维度上。二是BN的running_mean和running_var必须同步裁剪不能只剪weight——否则前面剪掉的通道索引错位模型前向计算出的特征全乱了后面微调也救不回来。另外这个实现假设每个Conv2d后面恰好跟着一个BN这也是2.1里刻意安排的结构约束。如果你的网络带残差连接ResNet的shortcut剪枝逻辑要复杂得多必须先对齐所有直接相加分支的通道再统一决定哪些通道剪掉否则相加操作两边通道数不匹配直接报错。VGG没有这个问题但如果你打算把这套代码往ResNet上迁移这一步必须重写。3.3 用FLOPs和文件大小双重验收剪完先别急着微调马上做两个检查。第一是用thop算参数量和FLOPs确认剪枝真的“瘦身”了第二是把state_dict存成文件看体积这才是对标“3M”目标的直接口径。我见过有人把剪枝写成mask置零非结构化参数量统计没变、模型文件也没变小还管它叫剪枝那只是把权重打碎离部署目标差得很远。from thop import profile pruned_model.eval() input_t torch.randn(1, 3, 32, 32) flops, params profile(pruned_model, inputs(input_t,)) print(f剪枝后 FLOPs: {flops / 1e6:.2f}M, Params: {params / 1e6:.2f}M) torch.save(pruned_model.state_dict(), vgg_pruned.pt) import os size_mb os.path.getsize(vgg_pruned.pt) / 1024 / 1024 print(f模型文件体积: {size_mb:.2f}MB)profile之前先把模型切到eval模式否则BN的统计量会在profile过程中被污染。如果你手头没有thop用fvcore的FlopCountAnalysis也可以两者计算FLOPs的口径略有差异但看趋势足够。这里要强调一个口径问题“3M”如果按文件体积算float32下约等于0.78M个参数如果按参数量算3M就是3百万参数对应约12MB的float32文件。标题里的3M在端侧部署语境里通常指文件体积本文后面全部按文件体积验收这样部署时flash占用才是可预期的。剪完第一次看文件体积时我预期你会落在8~20MB这个区间取决于稀疏强度而不是一步到3MB。不用慌真正的收尾在微调之后如果微调能稳住精度还可以进一步叠加半精度存储把体积再砍一半。那一步放在最后讲。4. 剪掉之后微调回来恢复精度的参数与循环策略4.1 剪完就评估精度暴跌多少才算正常刚剪完的模型不要直接拿去测试然后下结论说“剪枝没用”。先把剪完的模型在验证集上跑一遍分类准确率你大概率会看到一个从88%跌到70%甚至更低的数字。这是正常的原因有两个一是被剪掉的通道虽然gamma小但不是完全没贡献去掉后信息流变了二是BN的running_mean和running_var虽然被同步裁剪了但保留通道对应的统计量是在旧网络结构下统计的新网络里剩余通道的输出分布已经改变统计量失真。如果剪完精度直接掉到随机水平10%左右那不是剪枝策略的问题是索引对齐错了——大概率是上一章重建代码里的mask链断了或者BN统计量没裁剪。先回上一章查代码不要急着微调。如果掉10~20个点这是可恢复的区间微调可以拉回来大部分。我的经验是分成三个阶段看精度恢复微调5个epoch后应该回升到下降幅度的一半左右15~20个epoch后接近基线精度的90%以上30个epoch后趋于收敛。如果你微调了20个epoch精度还在低位徘徊说明剪枝率太激进回退到更保守的阈值重新剪而不是硬熬epoch数。4.2 微调超参模板学习率降到哪一档、训练多少轮微调不是从头训练权重已经处于一个较好的局部最优附近学率率太大会直接把权重踢出盆地。我一般把学习率降到基线训练的1/10到1/20也就是0.005左右配SGD训练30~40个epoch用余弦退火调度。weight_decay保持原来的一致不要再加L1稀疏项——微调阶段加L1会把好不容易保住的那部分通道又压扁这阶段的目标是让保留的通道重新适应新结构。微调的另一个选择是冻结部分浅层。前几层卷积提取的是边缘、颜色这类通用特征剪枝后它们的分布相对稳定可以对前两层设置requires_gradFalse只训练后半部分收敛更快。这个技巧在数据量大、微调预算紧张的时候尤其值得用。def finetune(model, train_loader, val_loader, epochs35, lr0.005): # 冻结前两层浅层特征相对通用先不动 for name, param in model.named_parameters(): if name.startswith(features.0) or name.startswith(features.3): param.requires_grad False optimizer torch.optim.SGD( filter(lambda p: p.requires_grad, model.parameters()), lrlr, momentum0.9, weight_decay5e-4 ) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_maxepochs) for ep in range(epochs): model.train() for images, targets in train_loader: optimizer.zero_grad() logits model(images) loss F.cross_entropy(logits, targets) loss.backward() optimizer.step() scheduler.step() # 每5个epoch顺手看验证精度 if (ep 1) % 5 0: model.eval() correct 0 with torch.no_grad(): for images, targets in val_loader: logits model(images) pred logits.argmax(dim1) correct (pred targets).sum().item() acc correct / len(val_loader.dataset) print(fepoch {ep1}: acc{acc:.4f}) return model冻结前两层的命名依赖features.0和features.3这两个具体索引如果你的cfg和我的不完全一致先打印model.features的层序号再写冻结逻辑。微调到后半程要盯验证集而不是训练集如果训练loss在降但验证精度不动说明新结构容量已经不够这时候加epoch没有意义该回头降低剪枝率。4.3 剪枝-微调循环逐步剪还是一步到位一步到位剪掉70%再微调是最快的方案但风险也最大。更稳的做法是剪枝-微调迭代循环每轮剪掉15~20%的通道剪完微调20个epoch等精度恢复到一个可接受的水平再剪下一轮。这样整个过程可能要多花2~3倍的训练时间但每一轮的网络结构变化小微调更容易把精度拉回来最终能在更高的整体剪枝率下保住质量。我的参考循环是这样组织初始全局阈值设成百分位30剪30%微调后记录验证精度下一轮在剩余通道上重新算gamma分布再剪20%微调重复直到精度降到不可接受或者模型体积达到目标。每一轮结束后都重新保存state_dict和mask因为下一轮剪枝的阈值必须基于微调后重新调整过的gamma计算而不是沿用上一轮的旧阈值。这个循环里最容易翻车的习惯是过度依赖验证集来决定一切。如果验证集本身不大反复在同一个验证集上挑“最好的检查点”最后得到的可能是过拟合验证集的剪枝率。我一般每轮固定微调epoch数只在全部循环结束后对比各轮留下的模型在单独测试集上的精度用测试集来验收而不是用验证集边调边测。5. 剪枝常见问题排查这5个坑每次都有人翻车5.1 剪枝后模型直接不收敛BN统计量与通道错位现象剪完后验证精度直接掉到10%附近微调20个epoch完全没有任何回升loss也降不下去。原因BN层在重建时只拷贝了weight和biasrunning_mean和running_var没有按mask裁剪或者裁剪时用错了索引。这两个缓存变量保存的是每个通道在训练期间统计的均值和方法通道索引错一位整个归一化输出就全乱了。另一种情况是重建时把第n层的输出mask用到了第n1层链式错位。解决把重建代码里涉及running_mean、running_var、weight、bias的裁剪索引统一用同一个out_keep变量不要分头写。如果已经剪坏了没有后悔药只能回到剪枝前重新加载稀疏训练好的模型再走一遍重建流程。我在本地总会保留一份剪枝前的state_dict快照这个习惯帮我省了很多重训时间。5.2 稀疏训练出来gamma一条线L1系数和warm-up在打架现象训练100个epoch后打印gamma分布发现所有gamma都在0.1到0.5之间没有明显的接近0的簇剪枝阈值无从下手。原因sparsity系数设小了L1惩罚在总loss里占比太低压不动gamma另一个可能是warm-up占了太长的周期比如前20个epoch都不加L1后半程稀疏项还没来得及发挥作用训练就结束了。解决把sparsity从1e-4提到5e-4同时把warm-up缩短到5个epoch以内。改成每10个epoch打印一次gamma的直方图或“小于0.01的比例”观察稀疏化是否在进行不要等100个epoch结束才看结果。这个监控习惯能让你在第30轮就发现参数配错省下大量重训时间。5.3 参数量降了模型文件却还很大旧state_dict与新结构不匹配现象用thop统计params从14M降到了3M但torch.save出来的文件体积还是40MB以上。原因你把剪枝后的model直接load了旧state_dict或者保存时用了整个model对象而不是state_dict。旧state_dict里每个参数的名字和shape都是按原始结构存的name匹配不上时load会静默失败模型用的其实是随机初始化的权重保存整个model对象则会把优化器状态、BN统计量缓存一并写入体积自然大。解决重建模型后先打印len(model.state_dict())和每个tensor的shape确认通道数确实收缩了再保存pruned_model.state_dict()。另外不要用torch.save(model)这种整对象保存应该保存state_dict加载时用VGGSlim(cfg)重建结构再load_state_dict。5.4 模型变小推理却更慢非结构化剪枝在GPU上的尴尬现象参数量少了70%部署到GPU上单张推理延迟反而比原模型高。原因如果剪枝只是把权重置零非结构化剪枝PyTorch和CUDA的常规矩阵乘不会跳过这些零元素计算量没变如果还用了稀疏张量存储CPU上有对应的稀疏算子还能快一点GPU上稀疏矩阵乘反而要额外的索引计算和显存访问除非用专为稀疏优化过的推理库否则只会更慢。剪枝的真正加速必须靠通道维度的收缩也就是让矩阵变“窄”而不是变“空”。解决检查你的剪枝是作用在结构上还是仅作用在数值上。本文实现的是通道剪枝剪完后的Conv2d的out_channels确实变小了FLOPs下降是真实的。如果追求进一步加速剪完后把模型转成ONNX或TensorRT再测端到端延迟不要在PyTorch的eager模式里死磕单次前向耗时。5.5 3M目标卡在瓶颈最后几MB藏在浮点精度里现象剪枝加微调后模型已经能到5~6MB但距3MB还差一截继续提高剪枝率则精度掉得太多卡住了。原因float32权重每个参数占4字节想从5MB降到3MB不只是通道数的事存储格式也是大头。如果剪枝率已经让精度逼近容忍下限继续“纯剪”性价比很低不如把存储格式从float32切到float16体积立刻减半4MB的模型降到2MB留出精度余量。解决在验证精度可接受的前提下把state_dict的权重转成half再保存torch.save({k: v.half() for k, v in model.state_dict().items()}, vgg_3m.pt)。加载推理时再把临时转回float32或者如果部署框架支持直接以half精度推理。这也是为什么3M这个目标通常被描述成“剪枝加量化”的组合动作单靠剪枝在结构上抠到3MB代价会非常难看。如果连half都试完还不够那就需要int8量化但验证流程就多一环放到下一章说。6. 验证与落盘用参数量、文件大小和耗时验收3M结果6.1 三分钟验证脚本精度、体积、延迟一次测完验收不能只看一个数字。剪枝任务的标准验收口径是三个指标同时看测试集精度、模型文件体积、单张图片推理耗时。只盯精度会忽略部署成本只盯体积又会剪过度。下面这个函数把三个指标打包剪完一轮就执行一次import time def evaluate_final(model, test_loader): model.eval() correct 0 with torch.no_grad(): for images, targets in test_loader: logits model(images) correct (logits.argmax(1) targets).sum().item() acc correct / len(test_loader.dataset) # 存state_dict并测体积 torch.save(model.state_dict(), vgg_final.pt) size_mb os.path.getsize(vgg_final.pt) / 1024 / 1024 # 测单张耗时取10次平均先warmup一次 sample torch.randn(1, 3, 32, 32) model(sample) # warmup start time.time() for _ in range(10): model(sample) latency_ms (time.time() - start) / 10 * 1000 print(facc{acc:.4f}, size{size_mb:.2f}MB, latency{latency_ms:.2f}ms) return acc, size_mb, latency_ms延迟这项在PyTorch的eager模式下存在较大波动CPU上多进程干扰、GPU上没有锁频都会影响读数所以只作为相对参考。真正部署时还要在ONNX Runtime或TensorRT里重新测一遍。如果这一步测出的体积离3MB还有距离优先检查是不是还有没剪干净的层比如最后一层Linear——它在GAP后本来就只有10个输出参数很少但不是0。6.2 继续压体积半精度与量化叠加如果float32剪枝后落在4~5MB先试half存储大概率能一步到3MB以内。half的效果很干净参数总量不变文件体积直接对半。代价是精度可能掉0.1~0.3个百分点对CIFAR-10的分类任务通常可以接受。如果half后精度掉得比你预期的多检查模型里是否有数值敏感的操作最常见的是最后一层Linear的权重值域分布不均匀可以只把这层保留float32其余层转half。int8量化是再往下一档的方案但不要无脑做。PyTorch的量化需要模型按量化感知训练的规范准备直接把剪枝后的模型feed给torch.quantization不一定能work因为VGGNet的BatchNorm和ReLU结构对量化范围的敏感性需要额外校准。我的建议是先把half作为3M目标的主力手段int8留到部署框架比如ONNX Runtime已经支持的情况下再去碰不要在剪枝流程里同时引入两个变量出了问题根本分不清是剪枝的问题还是量化的问题。6.3 我的验收习惯单卡复现清单最后把你需要准备的资源和预期结果盘一遍一张8GB显存的GPU或者性能好一点的CPU都可以跑完整流程数据集用CIFAR-10不用下载额外的大数据集PyTorch加thop两个依赖全流程耗时大约在4~6小时120 epoch稀疏训练约3小时后面每轮微调约30分钟。顺序是好结构GAP替换 → 稀疏训练120 epoch → 打印gamma分布确认稀疏化 → 通道剪枝 → 微调35 epoch → 三指标验收 → 不满意则调低剪枝率重来。我从这个流程里学到的一个教训是不要在一开始就把剪枝率定到70%以上先按50%走通全流程确认每个环节的数据都在预期轨道上再逐步收紧。第一次走通时建议把每步的中间模型都存一份因为等到第5章那些坑出现时任何一步的state_dict都可能成为你的后悔药。这套方法的价值不在于把VGGNet压到3M这个具体数字而在于训练、稀疏化、裁剪、恢复这个循环可以移植到任何带BN的卷积网络上。希望帮到你。本文还有配套的精品资源点击获取
返回列表