PyTorch中卷积层滤波器剪枝后如何更新预训练模型?
报错根因
出现维度不匹配错误的核心原因是卷积层剪枝后输出特征维度变化,没有同步修改后续全连接层的输入维度:
原第二层卷积输出通道为50,经过两次2×2池化后特征图尺寸为4×4,展平后维度是50×4×4=800;剪枝后第二层卷积输出通道变为45,展平后维度变成45×4×4=720,但全连接层第一个Linear的输入维度还是800,因此矩阵乘法无法执行。
手动更新剪枝模型的完整步骤
你需要同步调整卷积层和后续全连接层的参数,同时保留预训练权重中对应保留滤波器的部分:
- 调整卷积层参数
- 第一层卷积输出通道改为18,保留你筛选出的18个有效滤波器的权重和偏置
- 第二层卷积输入通道改为18(对应上一层输出通道)、输出通道改为45,保留你筛选出的45个有效滤波器的权重和偏置,且权重的输入通道维度要和上一层保留的滤波器索引对应
- 调整全连接层参数
- 第一层全连接层的输入维度改为720,截取原全连接层权重中对应720个有效输入维度的部分即可
参考代码如下:
- 第一层全连接层的输入维度改为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
相关产品推荐
相关产品推荐

