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

在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 11:43:11