为何torch.randn()初始化的nn.Parameter()未被模块注册?
问题根源与解决方法
核心差异与问题本质
- 第一种写法的问题出在调用顺序:
nn.Parameter(xxx).double()里的.double()是对已经创建好的Parameter实例调用的。因为nn.Parameter是torch.Tensor的子类,调用.double()会返回一个普通的DoubleTensor对象,不再是nn.Parameter类型。模块只会识别nn.Parameter实例作为可训练参数,所以这个普通张量不会被注册到参数列表里。 - 第二种写法是直接把
FloatTensor传入nn.Parameter的构造函数,返回的是标准的nn.Parameter实例,模块会自动将其注册为可训练参数,所以能正常出现在参数列表中。
正确的实现方式
你可以用两种方式修正第一种写法:
- 先将张量转换为double类型,再传入
nn.Parameter构造函数:
self.W = nn.Parameter(torch.randn(4,4).double())
- 先创建Parameter,再修改其数据类型(这种写法稍显繁琐,更推荐第一种):
self.W = nn.Parameter(torch.randn(4,4)) self.W.data = self.W.data.double()
另外补充:nn.Parameter的requires_grad参数默认就是True,所以完全可以省略这个参数,简化代码。
内容的提问来源于stack exchange,提问作者user4807817
相关产品推荐
相关产品推荐

