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

PyTorch中如何复用中间梯度,避免重复计算慢函数f的梯度

问题

假设我们有一个梯度计算缓慢的函数f,以及两个梯度计算简便的函数g1和g2。在PyTorch中,如何计算z1 = g1(f(x))和z2 = g2(f(x))关于x的梯度,同时避免重复计算函数f的梯度?

示例代码

import torch
import time

def slow_fun(x):
    A = x*torch.ones((1000,1000))
    B = torch.matrix_exp(1j*A)
    return torch.real(torch.trace(B))

x = torch.tensor(1.0, requires_grad = True)
y = slow_fun(x)
z1 = y**2
z2 = torch.sqrt(y)

start = time.time()
z1.backward(retain_graph = True)
end = time.time()
print("dz1/dx: ", x.grad)
print("duration: ", end-start, "\n")

x.grad = None
start = time.time()
z2.backward(retain_graph = True)
end = time.time()
print("dz2/dx: ", x.grad)
print("duration: ", end-start, "\n")

运行结果

dz1/dx:  tensor(-1673697.1250)
duration:  1.5571658611297607

dz2/dx:  tensor(-13.2334)
duration:  1.3989012241363525

可见计算dz2/dx的耗时与dz1/dx相近。如果PyTorch能在计算dz1/dx时存储dy/dx,并在计算dz2/dx时复用该结果,就能提升计算速度。请问PyTorch中是否有内置机制实现这一需求?


解决方案

PyTorch有内置机制可以实现这个需求,核心思路是预计算并复用f(x)对x的梯度dy/dx,再结合链式法则手动计算z1和z2的梯度,避免重复回溯f的计算图。具体有两种常用方式:

方法一:使用torch.autograd.grad预计算dy/dx

直接计算y对x的梯度并保存,之后通过链式法则(dz/dx = dz/dy * dy/dx)计算z1和z2的梯度:

import torch
import time

def slow_fun(x):
    A = x*torch.ones((1000,1000))
    B = torch.matrix_exp(1j*A)
    return torch.real(torch.trace(B))

x = torch.tensor(1.0, requires_grad = True)
y = slow_fun(x)
z1 = y**2
z2 = torch.sqrt(y)

# 预计算dy/dx,仅计算一次
start = time.time()
dy_dx, = torch.autograd.grad(y, x, retain_graph=True)
print("dy/dx计算耗时: ", time.time() - start, "\n")

# 计算dz1/dx:dz1/dy * dy/dx
start = time.time()
dz1_dy, = torch.autograd.grad(z1, y)
dz1_dx = dz1_dy * dy_dx
print("dz1/dx: ", dz1_dx)
print("dz1/dx计算耗时: ", time.time() - start, "\n")

# 计算dz2/dx:dz2/dy * dy/dx
start = time.time()
dz2_dy, = torch.autograd.grad(z2, y)
dz2_dx = dz2_dy * dy_dx
print("dz2/dx: ", dz2_dx)
print("dz2/dx计算耗时: ", time.time() - start, "\n")

这种方式下,slow_fun的梯度只计算一次,后续计算z1、z2的梯度时仅需快速计算g1、g2的梯度并相乘,大幅节省时间。

方法二:利用计算图保留与梯度累积

如果需要保留原有的backward调用逻辑,可以先计算dy/dx并保留计算图,之后通过向backward传递grad_tensors参数来复用已有的梯度信息:

import torch
import time

def slow_fun(x):
    A = x*torch.ones((1000,1000))
    B = torch.matrix_exp(1j*A)
    return torch.real(torch.trace(B))

x = torch.tensor(1.0, requires_grad = True)
y = slow_fun(x)
z1 = y**2
z2 = torch.sqrt(y)

# 先计算dy/dx并保留计算图
start = time.time()
y.backward(retain_graph=True)
dy_dx = x.grad.clone()
x.grad.zero_()
print("dy/dx计算耗时: ", time.time() - start, "\n")

# 计算dz1/dx:传递dz1/dy作为grad_tensors
start = time.time()
dz1_dy = 2*y  # 因为z1=y²,导数是2y
z1.backward(gradient=dz1_dy, retain_graph=True)
print("dz1/dx: ", x.grad)
print("dz1/dx计算耗时: ", time.time() - start, "\n")

x.grad.zero_()

# 计算dz2/dx:传递dz2/dy作为grad_tensors
start = time.time()
dz2_dy = 0.5 / torch.sqrt(y)  # z2=√y,导数是1/(2√y)
z2.backward(gradient=dz2_dy)
print("dz2/dx: ", x.grad)
print("dz2/dx计算耗时: ", time.time() - start, "\n")

这种方式同样只计算一次slow_fun的梯度,后续通过指定gradient参数直接复用已有的计算图节点,避免重复计算f的梯度。

关键原理

PyTorch的自动微分基于反向传播和计算图,当我们预计算dy/dx后,根据链式法则,z对x的梯度等于z对y的梯度乘以y对x的梯度。由于g1、g2的梯度计算很快,复用dy/dx就能避免重复执行slow_fun的反向传播过程,从而提升效率。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 06:27:18