Parameter 常见于 PyTorch 中,特别是在定义神经网络模型时。下面按常见场景给你说明 如何给 Parameter 赋值。
ParameterParameter 是 torch.nn.Parameter,本质是 Tensor,但会被自动注册为模型参数,参与优化和保存。
import torch
import torch.nn as nn
p = nn.Parameter(torch.tensor([1.0, 2.0]))
Parameter 赋值的几种方式import torch.nn as nn
class MyModel(nn.Module):
def __init__(self):
super().__init__()
self.weight = nn.Parameter(torch.randn(3, 3))
data)import torch
import torch.nn as nn
class MyModel(nn.Module):
def __init__(self):
super().__init__()
self.weight = nn.Parameter(torch.empty(3, 3))
# 赋值
self.weight.data = torch.randn(3, 3)
✅ 原因:
data 不会破坏计算图copy_ / fill_ / zero_ 等原地操作self.weight.data.fill_(1.0) # 全部赋值为 1
self.weight.data.copy_(tensor) # 拷贝另一个 tensor
self.weight.data.zero_() # 置 0
forward 中赋值(⚠不推荐)self.weight.data = torch.randn(3, 3)
❌ 问题:
self.weight = torch.randn(3, 3) # ❌ 不再是 Parameter
new_parameter = ...self.weight = nn.Parameter(torch.randn(3, 3)) # 会重新注册参数
import torch
import torch.nn as nn
class MyModel(nn.Module):
def __init__(self):
super().__init__()
self.weight = nn.Parameter(torch.empty(3, 3))
def init_weight(self):
self.weight.data = torch.eye(3)
model = MyModel()
model.init_weight()
print(model.weight)
| 场景 | 推荐方式 |
|---|---|
| 初始化参数 | nn.Parameter(tensor) |
| 后续赋值 | param.data = ... |
| 批量赋值 | param.data.copy_(tensor) |
| 避免 | 直接 self.param = tensor |
如果你指的是 其他框架(如 TensorFlow、C++、Java) 或 具体报错信息,可以贴出来,我帮你具体分析。
免责声明:本站发布的内容(图片、视频和文字)以原创、转载和分享为主,文章观点不代表本网站立场,如果涉及侵权请联系站长邮箱:is@yisu.com进行举报,并提供相关证据,一经查实,将立刻删除涉嫌侵权内容。