You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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文件尺寸与加载的原模型完全一致,并未如预期变小。我怀疑剪枝操作未生效,想请教出现该问题的原因及我的错误之处,谢谢!


问题原因与修复方案

关键问题分析

  1. 错误包含BatchNorm2d层剪枝
    PyTorch剪枝对nn.BatchNorm2d的权重仅做置零处理,不会移除任何参数张量的元素或维度——因为BatchNorm的权重是单通道标量,剪枝后参数总量不变。你将这类层加入剪枝列表,导致卷积/全连接层的剪枝效果被抵消,整体参数总量几乎无变化。

  2. 剪枝参数传递错误
    你在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.26 06:54:54