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

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}")

关键修改点

  1. 移除冗余的inp_copy和retain_graph=True参数
  2. 替换原地修改操作+=为新张量赋值,避免破坏反向传播所需的计算图
  3. 每轮迭代后用detach()更新inp,断开和当前轮计算图的关联,防止多轮计算图累积
  4. 损失函数替换为L1损失,避免原损失出现负值导致优化方向错误

扩展场景:同时优化输入inp和网络参数

如果需要将inp也作为可优化参数更新,仅需修改两处即可:

  1. 定义inp时开启梯度:inp = torch.tensor([1.0], requires_grad=True)
  2. 优化器添加inp作为参数:opt_updater = optim.Adam(list(net_updater.parameters()) + [inp], lr=1e-3)
  3. 移除迭代末尾的detach()操作,由优化器自动更新inp的值

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.26 05:15:05