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

