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
相关产品推荐
相关产品推荐

