PyTorch线性模型训练疑问:自定义参数与SGD优化器适配问题
PyTorch线性模型参数更新问题及结果不符分析
核心问题
练习用PyTorch实现简单线性模型,初始化了输入x、标签y以及自定义参数a、b,同时基于nn.Module创建了包含nn.Linear层的LinearModel,采用MSELoss作为损失函数。将自定义的a、b传入torch.optim.SGD优化器,但训练后模型结果与sklearn线性回归不符,疑惑这种参数传入方式是否能让模型参数在训练中更新,还是必须将参数嵌入model.parameters()中。
完整代码
初始化代码
import torch torch.manual_seed(42) x = torch.rand(100,1) y = 1 + 2 * x + .1 * torch.rand(100, 1) a = torch.randn(1, requires_grad=True) b = torch.randn(1, requires_grad=True)
模型创建代码
import torch.nn as nn from torch.optim import SGD # 创建简单线性模型 class LinearModel(nn.Module): def __init__(self): super().__init__() self.linear = nn.Linear(1, 1) def forward(self, x): return self.linear(x) # 损失函数 criterion = torch.nn.MSELoss() SGDmodel = LinearModel()
优化器代码
sgd = torch.optim.SGD([a, b], lr=0.001, momentum=0.9, weight_decay=0.1)
训练代码
epochs = 1 inputs = x targets = y for epoch in range(epochs): sgd.zero_grad() yhat = SGDmodel(inputs) loss = criterion(yhat, targets) loss.backward() sgd.step() print(f'Loss: {loss}')
梯度输出
print(SGDmodel.linear.weight.grad) print(SGDmodel.linear.bias.grad) # 输出: # tensor([[-3.1456]]) # tensor([-5.4119])
问题根源分析
当前代码存在核心逻辑错误:
- 自定义的
a、b与模型中nn.Linear层的weight、bias是完全独立的两组参数。 - 训练时,模型用
SGDmodel(inputs)计算预测值,本质是使用nn.Linear层自带的参数进行运算,反向传播后梯度会存储在SGDmodel.linear.weight.grad和SGDmodel.linear.bias.grad中。 - 但优化器
sgd绑定的是自定义的a、b,调用sgd.step()只会更新a、b的值,完全不会触及模型的实际参数,导致模型始终用初始随机参数预测,自然和sklearn的正确回归结果不符。
正确解决方案
有两种可行的修正方式,选择其一即可:
方案1:使用模型自带的参数(推荐)
直接让优化器管理模型nn.Linear层的参数,将优化器的参数列表改为SGDmodel.parameters():
sgd = torch.optim.SGD(SGDmodel.parameters(), lr=0.001, momentum=0.9, weight_decay=0.1)
这样训练时,反向传播产生的梯度会作用于模型的weight和bias,sgd.step()会正确更新这些参数,模型最终会拟合数据,结果也会和sklearn逐渐趋近。
方案2:手动定义模型参数并纳入管理
如果要使用自定义的a、b作为模型参数,需要将其注册为模型的属性,确保model.parameters()能包含它们,同时修改forward方法用这些参数计算:
class LinearModel(nn.Module): def __init__(self): super().__init__() # 将a、b注册为模型可训练参数 self.a = nn.Parameter(torch.randn(1)) self.b = nn.Parameter(torch.randn(1)) def forward(self, x): return self.a * x + self.b # 优化器传入模型参数 sgd = torch.optim.SGD(SGDmodel.parameters(), lr=0.001, momentum=0.9, weight_decay=0.1)
注意这里要用nn.Parameter包装参数,PyTorch会自动将其识别为模型的可训练参数,纳入参数管理体系。
内容的提问来源于stack exchange,提问作者Oscy
相关产品推荐
相关产品推荐

