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会直接报错。
问题根源
- 多次重复保留计算图:三次独立调用
torch.autograd.grad均设置retain_graph=True,导致计算图被重复保留,内存中累积大量冗余中间张量。 - 不必要的计算图保留:计算二阶导数时,
phi_xx和phi_yy均设置retain_graph=True,但实际仅需在第一个二阶导数计算时保留,最后一个可直接释放计算图。 - 冗余的张量操作:多次使用
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
相关产品推荐
相关产品推荐

