PyTorch训练时更新输入变量的正确方法及相关报错如何解决
报错原因解析
- 第一个inplace操作报错:你使用
inp += a原地修改了参与计算图的张量,同时加了retain_graph=True保留了上一轮的计算图,下一轮迭代修改inp后,上一轮保留的计算图中张量版本不匹配,触发报错。 - 第二个无梯度报错:
inp_copy是通过detach()生成的,该操作会将张量从计算图中剥离,不会携带梯度信息,基于它计算的loss自然无法反向传播更新网络参数。 - 第三个二次反向报错:你没有在每轮迭代结束后将inp从计算图中剥离,多轮迭代的计算图会不断叠加,第二轮反向传播时会尝试传播到已经被释放的第一轮计算图节点,触发该报错,并非显式调用了两次
backward。
正确实现方案
核心逻辑是每轮用网络生成更新量a,更新输入inp,同时优化网络参数让inp逐步接近目标值5,完整可运行代码如下:
import torch import torch.nn as nn import torch.optim as optim from torch.distributions import Normal class Model_updater(nn.Module): def __init__(self): super(Model_updater, self).__init__() self.fc1 = nn.Linear(1, 2) self.fc2 = nn.Linear(2, 3) self.fc3 = nn.Linear(3, 2) def forward(self, x): x = torch.relu(self.fc1(x)) x = torch.relu(self.fc2(x)) x = self.fc3(x) return x net_updater = Model_updater() # 可根据需求调整学习率 opt_updater = optim.Adam(net_updater.parameters(), lr=1e-3) inp = torch.tensor([1.0]) epochs = 100 for i in range(epochs): opt_updater.zero_grad() mu, sigma = net_updater(inp) dist1 = Normal(mu, torch.abs(sigma)) a = dist1.rsample() # 避免原地修改,生成新张量 inp_new = inp + a # 替换为L1损失,保证优化方向正确 loss = torch.abs(torch.tensor(5.0) - inp_new) loss.backward() opt_updater.step() # 断开当前轮计算图关联,避免多轮计算图叠加 inp = inp_new.detach() if i % 10 == 0: print(f"Epoch {i:>3d} | inp值: {inp.item():.4f} | loss: {loss.item():.4f}")
关键修改点
- 移除冗余的
inp_copy和retain_graph=True参数 - 替换原地修改操作
+=为新张量赋值,避免破坏反向传播所需的计算图 - 每轮迭代后用
detach()更新inp,断开和当前轮计算图的关联,防止多轮计算图累积 - 损失函数替换为L1损失,避免原损失出现负值导致优化方向错误
扩展场景:同时优化输入inp和网络参数
如果需要将inp也作为可优化参数更新,仅需修改两处即可:
- 定义inp时开启梯度:
inp = torch.tensor([1.0], requires_grad=True) - 优化器添加inp作为参数:
opt_updater = optim.Adam(list(net_updater.parameters()) + [inp], lr=1e-3) - 移除迭代末尾的
detach()操作,由优化器自动更新inp的值
内容的提问来源于stack exchange,提问作者Penguin
相关产品推荐
相关产品推荐

