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

