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

PyTorch中FNN计算二阶梯度内存上涨问题及优化咨询

PyTorch高阶梯度计算的内存膨胀问题解决与最佳实践

问题背景

在PyTorch中使用前馈神经网络(FNN)计算二阶导数时,出现迭代过程中内存占用持续上升的问题,即使调用gc.collect()和torch.cuda.empty_cache()也无法有效释放内存。相同架构在TensorFlow中未出现此现象。

问题代码简化版

def compute_gradients(self, x, y, t, x_2D=None):
    x = x.clone().detach().requires_grad_(True)
    y = y.clone().detach().requires_grad_(True)
    t = t.clone().detach().requires_grad_(True)

    g = torch.cat((x, y, t), dim=1)
    phi = self.ANN(g) 
    
    with torch.no_grad():
        phi_t = torch.autograd.grad(outputs=phi, inputs=t, grad_outputs=torch.ones_like(phi),
                                    create_graph=True, retain_graph=True)[0]
        phi_x = torch.autograd.grad(outputs=phi, inputs=x, grad_outputs=torch.ones_like(phi),
                                    create_graph=True, retain_graph=True)[0]
        phi_y = torch.autograd.grad(outputs=phi, inputs=y, grad_outputs=torch.ones_like(phi),
                                    create_graph=True, retain_graph=True)[0]

    with torch.no_grad():
        phi_xx = torch.autograd.grad(outputs=phi_x, inputs=x, grad_outputs=torch.ones_like(phi_x),
                                    create_graph=False, retain_graph=True)[0].detach()
        phi_yy = torch.autograd.grad(outputs=phi_y, inputs=y, grad_outputs=torch.ones_like(phi_y),
                                    create_graph=False, retain_graph=True)[0].detach()
    
    lap_phi = phi_xx + phi_yy
    lap_phi = lap_phi.detach().clone().requires_grad_(True)
    phi_t = phi_t.detach().clone().requires_grad_(True)
    
    del phi_x, phi_y, phi_xx, phi_yy, x, y, t, g
    torch.cuda.empty_cache()

    if x_2D is not None: 
        return phi, phi_t, lap_phi, phi2D_ann  
    else:
        return phi, phi_t, lap_phi

内存分析结果

Step: Memory usage In FNN, RSS: 57554.89 MB, VMS: 78540.54 MB
Step: Memory usage before phi_t, RSS: 57554.89 MB, VMS: 78540.54 MB
Step: Memory usage before phi_x, RSS: 57803.39 MB, VMS: 78788.55 MB
Step: Memory usage after phi_yy, RSS: 58371.66 MB, VMS: 79356.55 MB
Step: Memory usage Out FNN, RSS: 57579.65 MB, VMS: 78564.54 MB

可见每次计算二阶导数后内存显著上涨,将phi_xx和phi_yy的retain_graph设为False会直接报错。


问题根源

  1. 多次重复保留计算图:三次独立调用torch.autograd.grad均设置retain_graph=True,导致计算图被重复保留,内存中累积大量冗余中间张量。
  2. 不必要的计算图保留:计算二阶导数时,phi_xx和phi_yy均设置retain_graph=True,但实际仅需在第一个二阶导数计算时保留,最后一个可直接释放计算图。
  3. 冗余的张量操作:多次使用detach().clone().requires_grad_(True)创建额外张量副本,增加内存占用。

解决方案

1. 优化梯度计算逻辑,合并梯度调用

将多个一阶导数的计算合并为一次torch.autograd.grad调用,减少计算图的重复保留:

def compute_gradients(self, x, y, t, x_2D=None):
    # 仅在需要时克隆输入,避免冗余操作
    x = x.clone().detach().requires_grad_(True)
    y = y.clone().detach().requires_grad_(True)
    t = t.clone().detach().requires_grad_(True)

    g = torch.cat((x, y, t), dim=1)
    phi = self.ANN(g) 

    # 合并一阶梯度计算,仅保留一次计算图
    phi_t, phi_x, phi_y = torch.autograd.grad(
        outputs=phi,
        inputs=[t, x, y],
        grad_outputs=[torch.ones_like(phi)] * 3,
        create_graph=True,
        retain_graph=True  # 保留计算图用于二阶导数计算
    )

    # 计算二阶导数,仅在第一个时保留计算图,最后一个释放
    phi_xx = torch.autograd.grad(
        outputs=phi_x,
        inputs=x,
        grad_outputs=torch.ones_like(phi_x),
        create_graph=False,
        retain_graph=True  # 保留计算图用于phi_yy计算
    )[0].detach()
    
    phi_yy = torch.autograd.grad(
        outputs=phi_y,
        inputs=y,
        grad_outputs=torch.ones_like(phi_y),
        create_graph=False,
        retain_graph=False  # 最后一次计算,释放计算图
    )[0].detach()
    
    lap_phi = phi_xx + phi_yy
    # 若后续需要对lap_phi和phi_t求导,直接设置requires_grad=True,无需重新克隆
    lap_phi.requires_grad_(True)
    phi_t.requires_grad_(True)
    
    # 删除临时变量,触发Python垃圾回收
    del phi_x, phi_y, phi_xx, phi_yy, x, y, t, g
    import gc
    gc.collect()
    torch.cuda.empty_cache()

    if x_2D is not None: 
        return phi, phi_t, lap_phi, phi2D_ann  
    else:
        return phi, phi_t, lap_phi

2. 移除不必要的torch.no_grad()包裹

计算一阶导数时设置了create_graph=True,此时torch.no_grad()不会阻止计算图的创建,反而容易造成逻辑混淆,直接移除即可。


高阶梯度计算的内存管理最佳实践

  • 合并梯度计算:尽量通过一次torch.autograd.grad调用计算多个梯度,避免多次保留计算图导致的内存冗余。
  • 精准控制retain_graph:仅在后续还需使用当前计算图时设置retain_graph=True,最后一次梯度计算时设为False,及时释放计算图内存。
  • 避免冗余张量操作:若仅需修改张量的requires_grad属性,直接调用requires_grad_(True),无需先detach再clone。
  • 使用梯度检查点:对于大型模型,使用torch.utils.checkpoint.checkpoint以时间换空间,在反向传播时重新计算部分中间结果,减少内存占用。
  • 及时触发垃圾回收:在删除临时变量后,主动调用gc.collect()配合torch.cuda.empty_cache(),确保Python和CUDA的内存被及时释放。
  • 避免循环内累积计算图:在迭代训练时,确保每次迭代后计算图被完全释放,可将非梯度计算部分包裹在torch.no_grad()中。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.18 05:05:08