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

PyTorch自定义仅前向传播训练报错:重复反向遍历计算图

PyTorch自定义前向训练的计算图重复遍历错误解决

问题描述

尝试用PyTorch实现无需反向传播的自定义前向训练算法,手动更新神经网络权重,但首次迭代后反复出现“尝试第二次反向遍历计算图”的错误,即使执行了model.zero_grad()仍无法解决。

错误根源

  1. 计算图残留:初始化和step中调用self.f_x.backward()后,self.f_x、self.X、权重张量等仍保留计算图引用,第二次调用backward()时会尝试重新遍历已释放的旧图,引发错误。
  2. self.X的梯度追踪未断开:self.X初始设置了requires_grad=True,且更新时直接执行self.X = self.X + next_dX,新X仍关联旧计算图,导致后续计算的梯度与旧图纠缠。
  3. 梯度获取方式不当: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])

关键修改点说明

  1. self.X初始化与更新:初始化时移除requires_grad=True,更新时用detach()断开计算图,避免X关联旧训练计算图。
  2. 梯度获取方式:用torch.autograd.grad替代backward(),直接获取所需梯度,同时通过create_graph=True保证后续梯度计算可追踪。
  3. 移除冗余操作:删除step_theta中冗余的model.zero_grad(),避免混淆。
  4. 推理阶段优化:训练循环中调用forward时用torch.no_grad(),减少不必要的计算图开销。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 08:52:03