如何查看PyTorch网络各层参数及对应所属属性?
问题解答
问题1:未手动声明nn.Parameter也能返回参数的原因,以及查看参数的其他方法
- 原因:你调用的
nn.Linear是PyTorch内置的网络层类,它的内部已经自动将权重、偏置定义为nn.Parameter类型。同时,当你把nn.Linear这类继承自nn.Module的子模块赋值给自定义网络的实例属性时,父类nn.Module会自动递归收集所有子模块的参数,所以直接调用net.parameters()就能拿到全部参数。 - 查看参数的方法不止
.parameters()一种,常见的还有:.named_parameters():返回迭代器,每个元素是(参数名,参数值)的二元组,可以直接看到参数归属.state_dict():返回包含网络全部参数的有序字典,键为参数名,值为参数张量
问题2:参数顺序的对应关系,以及确认参数归属的方法
- 默认
.parameters()返回的参数顺序确实符合你说的规律:顺序为self.linear1权重 →self.linear1偏置 →self.linear2权重 →self.linear2偏置。这个顺序是按照你在__init__方法中定义子模块的先后,以及每个子模块内部参数定义的先后排列的。 - 要直接确认参数对应属性,直接调用
.named_parameters()遍历即可,示例代码如下:
for name, param in net.named_parameters(): print(f"参数名:{name},参数形状:{param.shape}")
运行后输出结果为:
参数名:linear1.weight,参数形状:torch.Size([2, 5]) 参数名:linear1.bias,参数形状:torch.Size([2]) 参数名:linear2.weight,参数形状:torch.Size([3, 2]) 参数名:linear2.bias,参数形状:torch.Size([3])
直接通过参数名就能直观对应到你定义的linear1、linear2的权重和偏置属性。
内容的提问来源于stack exchange,提问作者JAEMTO
相关产品推荐
相关产品推荐

