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

自定义Linear层网络报错RuntimeError:二次反向传播问题求助

PyTorch RuntimeError 解决方法

报错原因

你把self.xt设为模型的成员变量,每次forward时用它参与计算并更新,导致self.xt绑定了上一次前向传播的计算图节点。第二次调用loss.backward()时,PyTorch尝试回溯已经被释放的计算图(默认backward()执行后会释放计算图节省内存),因此抛出该错误。

解决方案

不要把循环状态xt作为模型成员变量,而是将其作为forward方法的输入参数,同时在forward中返回更新后的xt,让每次前向传播的计算图独立,避免跨batch的计算图关联。

修改后的代码

import torch
import torch.nn as nn

N = 10  # 根据你的实际场景调整N值

class MyNet(nn.Module):
    def __init__(self, n=N):
        super(MyNet, self).__init__()
        self.lx = nn.Linear(n, n)
        self.l1 = nn.Linear(n, n)
        self.l1_t = nn.Linear(n, n)
        self.ly = nn.Linear(n, 1)
        self.fn = nn.Tanh()

    def forward(self, x, xt):
        lx = self.fn(self.lx(x))
        l1 = self.fn(self.l1(lx + xt))

        updated_xt = self.fn(self.l1_t(xt + l1))

        output = self.fn(self.ly(l1))
        return output, updated_xt

使用示例

model = MyNet()
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)

# 初始化循环状态xt,维度与输入x匹配
xt = torch.zeros(N)

for epoch in range(10):
    # 模拟输入batch
    x = torch.randn(32, N)
    optimizer.zero_grad()
    
    # 传入当前xt,获取输出和更新后的状态
    output, xt = model(x, xt)
    
    # 计算损失并反向传播
    target = torch.randn(32, 1)
    loss = nn.MSELoss()(output, target)
    loss.backward()
    optimizer.step()

    # 可选:若需截断梯度防止爆炸,可添加xt = xt.detach()

额外说明

如果你的场景是类似RNN的循环状态传递,这种“状态作为输入输出”的方式是标准实现,能有效避免计算图复用导致的错误。若需要长期保留状态且不希望梯度累积,可在每次迭代后对xt调用detach()截断梯度。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 22:04:58