在PyTorch中如何对CIFAR10训练的CNN的MLP部分执行全局结构化剪枝?
全局结构化剪枝MLP模块的实现方案(PyTorch)
问题背景
你在训练CIFAR10的SimpleCNN模型时,希望对MLP模块(fc1、fc2、fc3)执行全局结构化剪枝:剪掉MLP中80%的神经元及其所有相关连接,而非对每层单独应用80%的剪枝比例。PyTorch内置的prune.ln_structured仅支持逐层剪枝,无法直接实现全局剪枝需求。
现成函数说明
PyTorch官方torch.nn.utils.prune模块目前没有直接支持全局结构化剪枝的函数,需要手动实现。
手动实现步骤
以下是基于神经元L2范数(与ln_structured一致的度量标准)的全局剪枝实现方案,核心思路是:先计算所有神经元的重要性,按全局比例筛选要保留的神经元,再重新构建MLP层以移除被剪神经元的连接。
1. 定义模型并获取MLP层
import torch import torch.nn as nn class SimpleCNN(nn.Module): def __init__(self): super().__init__() self.conv1 = nn.Conv2d(3, 32, kernel_size=3, padding=1) self.conv2 = nn.Conv2d(32, 64, kernel_size=3, padding=1) self.pool = nn.MaxPool2d(2, 2) self.relu = nn.ReLU() self.fc1 = nn.Linear(64 * 8 * 8, 128, bias=False) self.fc2 = nn.Linear(128, 128, bias=False) self.fc3 = nn.Linear(128, 10, bias=False) def forward(self, x): x = self.pool(self.relu(self.conv1(x))) x = self.pool(self.relu(self.conv2(x))) x = x.view(-1, 64 * 8 * 8) x = self.relu(self.fc1(x)) x = self.relu(self.fc2(x)) x = self.fc3(x) return x # 初始化模型 model = SimpleCNN() # 提取MLP层 mlp_layers = [model.fc1, model.fc2, model.fc3]
2. 计算所有神经元的重要性(L2范数)
对于每个全连接层,输出神经元的重要性用其对应权重行的L2范数表示:
neuron_scores = [] for layer_idx, layer in enumerate(mlp_layers): # 全连接层权重形状:[out_features, in_features] weights = layer.weight.data # 计算每个输出神经元权重行的L2范数 l2_norms = torch.norm(weights, p=2, dim=1) # 记录分数、所属层索引、神经元索引 for neuron_idx, score in enumerate(l2_norms): neuron_scores.append((score.item(), layer_idx, neuron_idx))
3. 确定全局剪枝比例并筛选要保留的神经元
这里以剪掉80%神经元为例,若不想剪输出层fc3,可调整总神经元数的计算范围:
# 计算总神经元数(包含fc3的输出神经元) total_neurons = sum(layer.out_features for layer in mlp_layers) # 若仅剪隐藏层,替换为: # total_neurons = mlp_layers[0].out_features + mlp_layers[1].out_features pruning_ratio = 0.8 # 计算要剪掉的神经元数量 num_to_prune = int(total_neurons * pruning_ratio) # 按分数升序排序(保留分数高的神经元) neuron_scores.sort(key=lambda x: x[0]) # 提取要剪掉的神经元 to_prune = neuron_scores[:num_to_prune]
4. 确定各层要保留的神经元索引
处理层间依赖:剪掉前层神经元后,后层的对应输入连接也需移除:
# 初始化各层保留的神经元集合 n1_keep = set(range(mlp_layers[0].out_features)) n2_keep = set(range(mlp_layers[1].out_features)) n3_keep = set(range(mlp_layers[2].out_features)) # 移除要剪的神经元 for _, layer_idx, neuron_idx in to_prune: if layer_idx == 0: n1_keep.remove(neuron_idx) elif layer_idx == 1: n2_keep.remove(neuron_idx) elif layer_idx == 2: n3_keep.remove(neuron_idx) # 转为有序列表 n1_keep = sorted(n1_keep) n2_keep = sorted(n2_keep) n3_keep = sorted(n3_keep)
5. 重新构建MLP层并替换原模型
# 重新构建fc1:保留指定输出神经元的权重行 new_fc1 = nn.Linear(mlp_layers[0].in_features, len(n1_keep), bias=False) new_fc1.weight.data = mlp_layers[0].weight.data[n1_keep, :].clone() # 重新构建fc2:保留指定输出神经元的权重行,以及对应fc1保留神经元的权重列 new_fc2 = nn.Linear(len(n1_keep), len(n2_keep), bias=False) new_fc2.weight.data = mlp_layers[1].weight.data[n2_keep, :][:, n1_keep].clone() # 重新构建fc3:保留指定输出神经元的权重行,以及对应fc2保留神经元的权重列 new_fc3 = nn.Linear(len(n2_keep), len(n3_keep), bias=False) new_fc3.weight.data = mlp_layers[2].weight.data[n3_keep, :][:, n2_keep].clone() # 替换原模型的MLP层 model.fc1 = new_fc1 model.fc2 = new_fc2 model.fc3 = new_fc3
补充说明
- 该方案直接修改模型层的维度,真正减少了参数数量和计算量,符合结构化剪枝的核心目标;
- 若仅需将权重置零(不改变层维度),可基于PyTorch的prune模块创建mask,但无法减少计算开销;
- 可根据需求替换神经元重要性的度量标准(如L1范数、激活值统计等)。
内容的提问来源于stack exchange,提问作者Noumeno
相关产品推荐
相关产品推荐

