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

PyTorch剪枝场景下,除module.weight/bias外如何访问权重与偏置?

PyTorch无结构化分层剪枝实现与参数访问疑问

1. 构建LeNet-300-100稠密神经网络

实现的LeNet-300-100模型代码如下:

import torch
import torch.nn as nn
import torch.nn.functional as F

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

class LeNet300(nn.Module):
    def __init__(self):
        super(LeNet300, self).__init__()
        
        # Define layers-
        self.fc1 = nn.Linear(in_features = 28 * 28 * 1, out_features = 300)
        self.fc2 = nn.Linear(in_features = 300, out_features = 100)
        self.output_layer = nn.Linear(in_features = 100, out_features = 10)
        
        self.weights_initialization()
    
    
    def forward(self, x):
        x = F.leaky_relu(self.fc1(x))
        x = F.leaky_relu(self.fc2(x))
        x = self.output_layer(x)
        return x
    
    
    def weights_initialization(self):
        '''
        When we define all the modules such as the layers in '__init__()'
        method above, these are all stored in 'self.modules()'.
        We go through each module one by one. This is the entire network,
        basically.
        '''
        for m in self.modules():
            if isinstance(m, nn.Linear):
                nn.init.kaiming_normal_(m.weight)
                nn.init.constant_(m.bias, 1)
    
    
    def shape_computation(self, x):
        print(f"Input shape: {x.shape}")
        x = self.fc1(x)
        print(f"dense1 output shape: {x.shape}")
        x = self.fc2(x)
        print(f"dense2 output shape: {x.shape}")
        x = self.output_layer(x)
        print(f"output shape: {x.shape}")
        del x
        return None

# Initialize architecture-
model = LeNet300().to(device)

该模型共有266610个可训练参数。

2. 剪枝需求与基础逻辑

剪枝要求:

  • 前两个全连接层(fc1、fc2)剪枝20%的权重
  • 输出层剪枝10%的权重
  • 经过25轮剪枝后达到99.5%的稀疏度
  • 使用torch.nn.utils.prune的l1_unstructured、random_unstructured方法实现无结构化分层剪枝

遍历网络层的代码:

for name, module in model.named_modules():
    if name == '':
        continue
    else:
        print(f"layer: {name}, module: {module}")

3. 分层剪枝实现思路与疑问

计划实现的分层剪枝代码(待加入偏置剪枝):

import torch.nn.utils.prune as prune

# Prune multiple parameters/layers in a given model-
for name, module in model.named_modules():
  
    # prune 20% of weights/connections in for all hidden layers-
    if isinstance(module, torch.nn.Linear) and name != 'output_layer':
        prune.l1_unstructured(module = module, name = 'weight', amount = 0.2)
    
    # prune 10% of weights/connections for output layer-
    elif isinstance(module, torch.nn.Linear) and name == 'output_layer':
        prune.l1_unstructured(module = module, name = 'weight', amount = 0.1)

疑问:除了直接使用module.weight、module.bias,还有哪些方法可以访问特定模块的权重与偏置?


访问特定模块参数的实用方法

1. 直接通过模型属性访问

因为模型类中已经显式定义了self.fc1、self.fc2、self.output_layer,可以直接通过模型实例的属性精准访问:

# 访问fc1的权重
fc1_weight = model.fc1.weight
# 访问fc2的偏置
fc2_bias = model.fc2.bias
# 访问输出层的权重
output_weight = model.output_layer.weight

这种方式最直接,适合明确知道目标模块名称的场景。

2. 通过named_parameters()遍历筛选

遍历模型所有可训练参数,通过参数名称匹配目标:

for param_name, param in model.named_parameters():
    if param_name == "fc1.weight":
        # 处理fc1的权重
        pass
    elif param_name == "output_layer.bias":
        # 处理输出层的偏置
        pass

参数名称的格式为模块名.参数名,比如fc1的偏置对应fc1.bias。

3. 结合named_modules()与getattr()动态获取

在遍历模块的循环中,用getattr()方法动态获取模块内的权重或偏置,适合批量处理场景:

for name, module in model.named_modules():
    if isinstance(module, torch.nn.Linear):
        # 动态获取权重
        weight = getattr(module, 'weight')
        # 动态获取偏置
        bias = getattr(module, 'bias')
        # 执行剪枝操作
        if name != 'output_layer':
            prune.l1_unstructured(module=module, name='weight', amount=0.2)
            prune.l1_unstructured(module=module, name='bias', amount=0.1)  # 偏置剪枝
        else:
            prune.l1_unstructured(module=module, name='weight', amount=0.1)
            prune.l1_unstructured(module=module, name='bias', amount=0.05)

4. 剪枝后参数的访问注意

使用torch.nn.utils.prune剪枝后,原参数会被重命名为weight_orig/bias_orig,而module.weight/module.bias是应用了剪枝掩码的参数。如果需要访问原始未剪枝的参数,用module.weight_orig;如果要获取剪枝掩码,用module.weight_mask。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.13 22:10:29