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

PyTorch中卷积层滤波器剪枝后如何更新预训练模型?

报错根因

出现维度不匹配错误的核心原因是卷积层剪枝后输出特征维度变化,没有同步修改后续全连接层的输入维度:
原第二层卷积输出通道为50,经过两次2×2池化后特征图尺寸为4×4,展平后维度是50×4×4=800;剪枝后第二层卷积输出通道变为45,展平后维度变成45×4×4=720,但全连接层第一个Linear的输入维度还是800,因此矩阵乘法无法执行。

手动更新剪枝模型的完整步骤

你需要同步调整卷积层和后续全连接层的参数,同时保留预训练权重中对应保留滤波器的部分:

  1. 调整卷积层参数
    • 第一层卷积输出通道改为18,保留你筛选出的18个有效滤波器的权重和偏置
    • 第二层卷积输入通道改为18(对应上一层输出通道)、输出通道改为45,保留你筛选出的45个有效滤波器的权重和偏置,且权重的输入通道维度要和上一层保留的滤波器索引对应
  2. 调整全连接层参数
    • 第一层全连接层的输入维度改为720,截取原全连接层权重中对应720个有效输入维度的部分即可
      参考代码如下:
import torch
import torch.nn as nn
import torch.nn.functional as F

# 以下变量需替换为你实际剪枝得到的保留滤波器索引
# 第一层卷积保留的18个滤波器的索引列表,长度为18
remaining_conv1_idx = [0,1,2,3,4,5,6,7,8,9,10,11,12,13,14,15,16,17] 
# 第二层卷积保留的45个滤波器的索引列表,长度为45
remaining_conv2_idx = [i for i in range(45)] 

# 构建剪枝后的新模型结构
class PrunedLeNet5(nn.Module):
    def __init__(self, n_classes):
        super().__init__()
        self.feature_extractor = nn.Sequential(            
            nn.Conv2d(in_channels=1, out_channels=18, kernel_size=5, stride=1),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=2),
            nn.Conv2d(in_channels=18, out_channels=45, kernel_size=5, stride=1),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=2)
        )
        self.classifier = nn.Sequential(
            nn.Linear(in_features=720, out_features=500),
            nn.ReLU(),
            nn.Linear(in_features=500, out_features=10),
        )
    
    def forward(self, x):
        x = self.feature_extractor(x)
        x = torch.flatten(x, 1)
        logits = self.classifier(x)
        probs = F.softmax(logits, dim=1)
        return logits, probs

# 初始化新模型
pruned_model = PrunedLeNet5(n_classes=10)
# 假设old_model是你训练好的原始LeNet5模型
# 赋值卷积层权重
pruned_model.feature_extractor[0].weight.data = old_model.feature_extractor[0].weight.data[remaining_conv1_idx, :, :, :]
pruned_model.feature_extractor[0].bias.data = old_model.feature_extractor[0].bias.data[remaining_conv1_idx]
pruned_model.feature_extractor[3].weight.data = old_model.feature_extractor[3].weight.data[remaining_conv2_idx, :, :, :][:, remaining_conv1_idx, :, :]
pruned_model.feature_extractor[3].bias.data = old_model.feature_extractor[3].bias.data[remaining_conv2_idx]
# 赋值全连接层权重
pruned_model.classifier[0].weight.data = old_model.classifier[0].weight.data[:, :720]
pruned_model.classifier[0].bias.data = old_model.classifier[0].bias.data
pruned_model.classifier[2].weight.data = old_model.classifier[2].weight.data
pruned_model.classifier[2].bias.data = old_model.classifier[2].bias.data

可用的剪枝工具库

PyTorch官方自带torch.nn.utils.prune工具,支持结构化、非结构化等多种剪枝方式,可自动处理层间参数匹配、权重掩码管理,不需要手动修改每层的参数维度。除此之外也可以使用TorchPrune等第三方剪枝库实现更复杂的剪枝策略,所有接口都兼容原生PyTorch逻辑。

内容的提问来源于stack exchange,提问作者NIKITA RATH

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 19:06:01