You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch技术问题:未知输入维度时如何创建并延迟初始化Parameter

动态初始化PyTorch Module参数的正确实现

现有代码的问题

  • 初始创建self.weight = torch.nn.Parameter(torch.FloatTensor(None, out_dim))不合法:PyTorch不支持带None维度的可训练参数,这类张量无法被正确注册为模型参数。
  • expand方法仅返回张量视图,不会分配新内存,无法作为可训练参数的有效初始化;且后续直接赋值self.weight会导致优化器丢失对该参数的跟踪——优化器是在模型初始化时获取的参数列表,新赋值的参数不在其中。
  • 初始化函数中的xvaier_normal_拼写错误,正确写法是xavier_normal_。

正确实现方案

方案1:延迟参数创建,同步调整优化器初始化时机

先不创建参数,在第一次forward时根据输入维度生成参数,之后再初始化优化器,确保优化器能捕获到所有可训练参数。

修改后的模型代码:

import torch

class MyModel(torch.nn.Module):
    def __init__(self, out_dim):
        super(MyModel, self).__init__()
        self.out_dim = out_dim
        self.weight = None
        self.init_weight = False
    
    def init_parameters(self, in_dim):
        # 创建指定维度的可训练参数
        self.weight = torch.nn.Parameter(torch.empty(in_dim, self.out_dim))
        # 用xavier方法初始化权重
        torch.nn.init.xavier_normal_(self.weight)
        self.init_weight = True

    def forward(self, X):
        if not self.init_weight:
            self.init_parameters(X.shape[1])
        
        return torch.sigmoid(torch.matmul(X, self.weight))        

修改后的训练代码:

def train(X_train, y_train):
    model = MyModel(y_train.shape[1])
    loss_fn = torch.nn.MSELoss()

    # 先执行一次forward触发参数初始化,无需计算梯度
    with torch.no_grad():
        model(X_train)
    
    # 参数已创建,此时初始化优化器
    optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)

    for epoch in range(10000):
        optimizer.zero_grad()

        prediction = model(X_train)
        loss = loss_fn(prediction, y_train)
        loss.backward()

        optimizer.step()

方案2:提前初始化优化器,手动添加新参数到优化器

如果必须提前初始化优化器,可以在参数创建后,手动将其加入优化器的参数组。

模型代码与方案1一致,训练代码修改为:

def train(X_train, y_train):
    model = MyModel(y_train.shape[1])
    optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
    loss_fn = torch.nn.MSELoss()

    for epoch in range(10000):
        optimizer.zero_grad()

        prediction = model(X_train)
        
        # 第一次初始化参数后,将其添加到优化器
        if epoch == 0 and model.init_weight:
            optimizer.add_param_group({'params': model.weight})

        loss = loss_fn(prediction, y_train)
        loss.backward()

        optimizer.step()

关键注意事项

  • 动态创建参数时,必须用torch.nn.Parameter包装,直接赋值self.weight会自动将参数注册到模型,也可以用self.register_parameter('weight', param)显式注册。
  • 优化器必须跟踪到新参数才能进行梯度更新,因此要么延迟优化器初始化到参数创建后,要么手动将新参数加入优化器参数组。
  • 权重初始化方法要注意拼写正确性,避免因笔误导致初始化失败。

内容的提问来源于stack exchange,提问作者LuoYuanzhen

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.21 18:35:23