PyTorch中ResNet18剪枝后模型尺寸未减小的问题排查
PyTorch ResNet18剪枝后模型文件大小未变化的问题
我尝试在PyTorch中对ResNet18模型实现剪枝操作,加载预训练模型后执行如下代码进行剪枝并保存模型:
def random_unstructured_pruning(pruning_rate: float, device, log_file): trained_model=retrieve_file(folder="./models", file_name='trained_model.pth') model=ResNet18(num_classes=10, input_channels=1).to(device) model.load_state_dict(torch.load(trained_model)) modules_list=filter(lambda x: isinstance(x[1], (nn.Conv2d, nn.Linear, nn.BatchNorm2d)), model.named_modules()) modules_list = map(lambda x: (x[1], 'weight'), modules_list) modules_list=tuple(modules_list) prune.global_unstructured(modules_list, pruning_method=prune.L1Unstructured, amount=0.8) for module in modules_list: prune.remove(module[0], module[1]) pruning_rate_str= "{:02d}".format(int(pruning_rate * 10)) path=f"{model_saving_path}pruned_{pruning_rate_str}.pth" # torch.save(model.state_dict(), f"{path}")
我设置了80%的剪枝率,但最终保存的.pth文件尺寸与加载的原模型完全一致,并未如预期变小。我怀疑剪枝操作未生效,想请教出现该问题的原因及我的错误之处,谢谢!
问题原因与修复方案
关键问题分析
错误包含BatchNorm2d层剪枝
PyTorch剪枝对nn.BatchNorm2d的权重仅做置零处理,不会移除任何参数张量的元素或维度——因为BatchNorm的权重是单通道标量,剪枝后参数总量不变。你将这类层加入剪枝列表,导致卷积/全连接层的剪枝效果被抵消,整体参数总量几乎无变化。剪枝参数传递错误
你在prune.global_unstructured中硬编码了amount=0.8,未使用函数传入的pruning_rate参数;同时,全局剪枝的amount是总参数的剪去比例,混入BatchNorm参数后,实际可压缩的卷积/全连接参数占比被稀释,剪枝效果微乎其微。
修复后的代码
def random_unstructured_pruning(pruning_rate: float, device, log_file): trained_model=retrieve_file(folder="./models", file_name='trained_model.pth') model=ResNet18(num_classes=10, input_channels=1).to(device) model.load_state_dict(torch.load(trained_model)) # 仅对Conv2d和Linear层的权重执行剪枝,排除BatchNorm2d modules_list=filter(lambda x: isinstance(x[1], (nn.Conv2d, nn.Linear)), model.named_modules()) modules_list = map(lambda x: (x[1], 'weight'), modules_list) modules_list=tuple(modules_list) # 使用传入的pruning_rate作为剪枝比例 prune.global_unstructured(modules_list, pruning_method=prune.L1Unstructured, amount=pruning_rate) for module in modules_list: prune.remove(module[0], module[1]) # 修正剪枝率格式化逻辑:*100才能正确显示百分比对应的数值(如80%→80) pruning_rate_str= "{:02d}".format(int(pruning_rate * 100)) path=f"{model_saving_path}pruned_{pruning_rate_str}.pth" torch.save(model.state_dict(), path)
验证剪枝效果的方法
剪枝前后可以打印参数总量对比,确认剪枝是否生效:
# 剪枝前统计参数总数 total_params_before = sum(p.numel() for p in model.parameters()) # 执行剪枝及prune.remove操作后 total_params_after = sum(p.numel() for p in model.parameters()) print(f"剪枝前参数总量: {total_params_before}, 剪枝后参数总量: {total_params_after}")
若参数总量明显减少,说明剪枝生效,保存后的模型文件大小会相应降低。
内容的提问来源于stack exchange,提问作者Alessandro Masci
相关产品推荐
相关产品推荐

