温馨提示×

温馨提示×

您好,登录后才能下订单哦!

密码登录×
登录注册×
其他方式登录
点击 登录注册 即表示同意《亿速云用户服务条款》

Parameter怎么赋值

发布时间:2026-07-23 17:10:43 来源:亿速云 阅读:84 作者:小樊 栏目:编程语言

Parameter 常见于 PyTorch 中,特别是在定义神经网络模型时。下面按常见场景给你说明 如何给 Parameter 赋值


一、什么是 Parameter

Parametertorch.nn.Parameter,本质是 Tensor,但会被自动注册为模型参数,参与优化和保存。

import torch
import torch.nn as nn

p = nn.Parameter(torch.tensor([1.0, 2.0]))

二、给 Parameter 赋值的几种方式

✅ 方法 1:创建时直接赋值(最常用)

import torch.nn as nn

class MyModel(nn.Module):
    def __init__(self):
        super().__init__()
        self.weight = nn.Parameter(torch.randn(3, 3))

✅ 方法 2:先定义,再赋值(推荐用 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 不会破坏计算图
  • 不会被当成新的 Parameter 重新注册

✅ 方法 3:用 copy_ / fill_ / zero_ 等原地操作

self.weight.data.fill_(1.0)        # 全部赋值为 1
self.weight.data.copy_(tensor)     # 拷贝另一个 tensor
self.weight.data.zero_()           # 置 0

✅ 方法 4:在 forward 中赋值(⚠不推荐)

self.weight.data = torch.randn(3, 3)

❌ 问题:

  • 每次 forward 都会改变参数
  • 影响梯度、优化器状态

三、错误示例(不要这样写)

❌ 直接赋值(会丢失 Parameter 属性)

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)具体报错信息,可以贴出来,我帮你具体分析。

向AI问一下细节

免责声明:本站发布的内容(图片、视频和文字)以原创、转载和分享为主,文章观点不代表本网站立场,如果涉及侵权请联系站长邮箱:is@yisu.com进行举报,并提供相关证据,一经查实,将立刻删除涉嫌侵权内容。

AI