如何在PyTorch中查看预训练EfficientNet-B3的层权重?
在PyTorch中查看EfficientNet-B3的层权重
问题背景
想要查看已加载预训练权重的EfficientNet-B3模型中层的权重,尝试了以下代码但未达到预期效果:
os.system('pip install efficientnet_pytorch') from efficientnet_pytorch import EfficientNet MODEL_NAME = 'efficientnet-b3' effnet = EfficientNet.from_pretrained(MODEL_NAME) effnet.modules # 仅返回模块名称相关信息 effnet.weights # 执行报错 effnet.layers # 执行报错 effnet.modules[1] # 执行报错,目标是获取第二个模块(批归一化层_bn0)的权重
希望实现类似TensorFlow代码的功能,获取第一个批归一化层的权重:
from tensorflow.keras.applications import EfficientNetB3 base_model = EfficientNetB3(weights="imagenet") base_model.trainable_variables[1].numpy() # 通过索引1获取批归一化层权重
上述TensorFlow代码的输出示例:
<tf.Variable 'stem_bn/gamma:0' shape=(40,) dtype=float32, numpy= array([ 0.1913209 , 2.7074034 , 9.623442 , 2.5562265 , 3.127593 , 4.348222 , 2.4381876 , 3.4623973 , 3.6115906 , 4.1241236 , 2.18851 , 8.9716835 , 0.7232651 , 0.6261555 , 9.050293 , 7.9233327 , 0.47725916, 3.4991856 , 5.334402 , 4.843143 , 1.4122163 , 1.953061 , 8.150878 , 5.0044165 , 2.3806598 , 4.2976685 , 2.2239766 , 0.551327 , 7.799995 , 3.3823645 , 1.8910869 , 4.0793633 , 0.73215246, 3.4526935 , 10.874565 , 2.0920732 , 6.272054 , 3.6823177 , 4.2152214 , 3.4319222 ], dtype=float32)>
解决方案
在PyTorch中,模型层的权重以Parameter对象存储,以下几种方法可以获取目标层的权重:
1. 直接访问指定子模块
EfficientNet-B3的第一个批归一化层对应模型的_bn0属性,直接访问即可获取权重(对应TensorFlow中的gamma)和偏置(对应beta):
import torch # 获取BN层的权重(gamma) bn_weight = effnet._bn0.weight.data.numpy() # 获取BN层的偏置(beta) bn_bias = effnet._bn0.bias.data.numpy()
运行后得到的bn_weight就是形状为(40,)的数组,和TensorFlow示例的输出一致。
2. 遍历模块查找目标层
如果不确定层的名称,可以通过named_modules()遍历所有模块,筛选出目标批归一化层:
import torch for name, module in effnet.named_modules(): if isinstance(module, torch.nn.BatchNorm2d) and name == '_bn0': print(f"模块名称: {name}") print(f"权重数组: {module.weight.data.numpy()}") print(f"偏置数组: {module.bias.data.numpy()}") break
3. 处理modules生成器
effnet.modules是一个生成器,不能直接用索引访问,需要先转换成列表:
modules_list = list(effnet.modules()) # 列表的第二个元素就是目标批归一化层 bn_module = modules_list[1] bn_weight = bn_module.weight.data.numpy()
4. 查看所有可训练参数
如果需要查看模型全部可训练参数,使用named_parameters()遍历:
for name, param in effnet.named_parameters(): print(f"参数名称: {name}") print(f"参数形状: {param.shape}") # 如需查看具体数值,添加 print(param.data.numpy())
内容的提问来源于stack exchange,提问作者Kilaru Vasudeva
相关产品推荐
相关产品推荐

