torch.nn.Parameter
基本介绍
torch.nn.Parameter是继承自torch.Tensor的子类,其主要作用是作为nn.Module中的可训练参数使用。它与torch.Tensor的区别就是nn.Parameter会自动被认为是module的可训练参数,即加入到parameter()这个迭代器中去。
具体格式如下:
torch.nn.parameter.Parameter(data=None, requires_grad=True)
其中 data 为待传入的 Tensor,requires_grad 默认为 True。
事实上,torch.nn 中提供的模块中的参数均是 nn.Parameter 类,例如:
module = nn.Linear(3, 3)
type(module.weight)
# torch.nn.parameter.Parameter
type(module.bias)
# torch.nn.parameter.Parameter
参数构造
nn.Parameter可以看作是一个类型转换函数,将一个不可训练的类型 Tensor 转换成可以训练的类型 parameter ,并将这个 parameter 绑定到这个module 里面nn.Parameter()添加的参数会被添加到Parameters列表中,会被送入优化器中随训练一起学习更新
此时调用 parameters()方法会显示参数。读者可自行体会以下两端代码:
""" 代码片段一 """
class Net(nn.Module):
def __init__(self):
super().__init__()
self.weight = torch.randn(3, 3)
self.bias = torch.randn(3)
def forward(self, inputs):
pass
net = Net()
print(list(net.parameters()))
# []
""" 代码片段二 """
class Net(nn.Module):
def __init__(self):<

本文详细介绍了PyTorch中torch.nn.Parameter的使用方法,包括参数构造、访问、初始化及参数绑定等内容。通过实例展示了如何利用nn.Parameter进行参数管理,以及如何进行参数初始化。

441

被折叠的 条评论
为什么被折叠?



