PyTorch自定义仅前向传播训练报错:重复反向遍历计算图
PyTorch自定义前向训练的计算图重复遍历错误解决
问题描述
尝试用PyTorch实现无需反向传播的自定义前向训练算法,手动更新神经网络权重,但首次迭代后反复出现“尝试第二次反向遍历计算图”的错误,即使执行了model.zero_grad()仍无法解决。
错误根源
- 计算图残留:初始化和
step中调用self.f_x.backward()后,self.f_x、self.X、权重张量等仍保留计算图引用,第二次调用backward()时会尝试重新遍历已释放的旧图,引发错误。 self.X的梯度追踪未断开:self.X初始设置了requires_grad=True,且更新时直接执行self.X = self.X + next_dX,新X仍关联旧计算图,导致后续计算的梯度与旧图纠缠。- 梯度获取方式不当:
backward()会自动累积梯度并保留计算图,手动训练场景下更适合用torch.autograd.grad直接获取梯度,避免计算图残留。
修复方案
- 用
torch.autograd.grad替代backward()获取梯度,默认不保留计算图。 - 更新
self.X时通过detach()断开计算图关联,确保新X不再追踪梯度。 - 手动更新权重时,确保所有参与计算的中间张量不保留梯度引用,避免计算图累积。
修改后的完整代码
模型类代码
import torch import numpy as np class OneD_NN_LQR: def __init__(self, hidden_units, learning_rate_param_C=0.05, batch_size=100): self.C = learning_rate_param_C self.N = batch_size self.hidden_units = hidden_units self.dim = 1 self.layer1 = torch.nn.Linear(in_features=self.dim, out_features=self.hidden_units) self.activation = torch.nn.ReLU() self.layer2 = torch.nn.Linear(in_features=self.hidden_units, out_features=self.dim, bias=False) self.model = torch.nn.Sequential( self.layer1, self.activation, self.layer2 ) self.w = self.layer1.weight self.b = self.layer1.bias self.c = self.layer2.weight self.Xtilde_w = torch.zeros((self.hidden_units,)).unsqueeze(1) self.Xtilde_c = torch.zeros((self.hidden_units,)).unsqueeze(1) self.Xtilde_b = torch.zeros((self.hidden_units,)) # 初始化X时无需追踪梯度(手动更新) self.X = torch.ones((self.dim,)) self.f_x = self.forward(self.X) # 用autograd.grad获取梯度,create_graph=True保证后续梯度计算可追踪 grads = torch.autograd.grad(self.f_x, inputs=[self.X, self.c, self.b, self.w], create_graph=True) self.grad_x = grads[0] self.grad_c = grads[1].T self.grad_b = grads[2] self.grad_w = grads[3] self.time = 0 def step(self, delta): self.time += delta self.step_X(delta) self.step_Xtilde(delta) self.step_theta(delta) self.model.zero_grad() # 重新计算f_x,此时X已断开旧计算图 self.f_x = self.forward(self.X) print(self.f_x) # 再次用autograd.grad获取梯度 grads = torch.autograd.grad(self.f_x, inputs=[self.X, self.c, self.b, self.w], create_graph=True) self.grad_x = grads[0] self.grad_c = grads[1].T self.grad_b = grads[2] self.grad_w = grads[3] return self.w, self.c, self.b def step_theta(self, delta): next_dw, next_dc, next_db = self.next_dtheta(delta) with torch.no_grad(): self.layer1.weight.sub_(next_dw) self.layer1.bias.sub_(next_db) self.layer2.weight.sub_(next_dc.T) def step_X(self, delta): next_dX = self.next_dX(delta) # 断开计算图,避免X关联旧图 self.X = (self.X + next_dX).detach() def step_Xtilde(self, delta): next_dXtilde_w, next_dXtilde_c, next_dXtilde_b = self.next_dXtilde(delta) self.Xtilde_w = self.Xtilde_w + next_dXtilde_w self.Xtilde_c = self.Xtilde_c + next_dXtilde_c self.Xtilde_b = self.Xtilde_b + next_dXtilde_b def next_dtheta(self, delta): alpha = self.get_learning_rate(self.C, self.time) dw = alpha * (2 * self.X * self.Xtilde_w + 2 * self.f_x * (self.grad_w + self.grad_x * self.Xtilde_w)) * delta db = alpha * (2 * self.X * self.Xtilde_b + 2 * self.f_x * (self.grad_b + self.grad_x * self.Xtilde_b)) * delta dc = alpha * (2 * self.X * self.Xtilde_w + 2 * self.f_x * (self.grad_c + self.grad_x * self.Xtilde_c)) * delta return dw, dc, db def get_learning_rate(self, c, time): if time > 500: return c / 10 if time > 100: return c / 5 if time > 50: return c / 2 return c def next_dXtilde(self, delta): dXtilde_w = (- self.Xtilde_w + self.grad_w + self.grad_x * self.Xtilde_w) * delta dXtilde_b = (- self.Xtilde_b + self.grad_b + self.grad_x * self.Xtilde_b) * delta dXtilde_c = (- self.Xtilde_c + self.grad_c + self.grad_x * self.Xtilde_c) * delta return dXtilde_w, dXtilde_c, dXtilde_b def next_dX(self, delta): to_return = (-self.X + self.f_x) * delta + torch.normal(0, 1, size=(self.dim,)) * (delta ** 0.5) return to_return def forward(self, x): to_return = self.model.forward(x) return to_return
训练循环代码
x = torch.tensor([5]).unsqueeze(1) y = [] step_size = 1e-2 theta_vals = [] range_end = 10 fwd_propagator = OneD_NN_LQR(16, learning_rate_param_C=100, batch_size=10) for i in np.arange(0, range_end, step_size): theta = fwd_propagator.step(step_size) theta_vals.append(theta) # 推理阶段禁用梯度追踪,减少开销 with torch.no_grad(): y.append(fwd_propagator.forward(x)[0])
关键修改点说明
self.X初始化与更新:初始化时移除requires_grad=True,更新时用detach()断开计算图,避免X关联旧训练计算图。- 梯度获取方式:用
torch.autograd.grad替代backward(),直接获取所需梯度,同时通过create_graph=True保证后续梯度计算可追踪。 - 移除冗余操作:删除
step_theta中冗余的model.zero_grad(),避免混淆。 - 推理阶段优化:训练循环中调用
forward时用torch.no_grad(),减少不必要的计算图开销。
内容的提问来源于stack exchange,提问作者Gummy bears
相关产品推荐
相关产品推荐

