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
相关产品推荐
相关产品推荐

