原地操作引发内存泄漏?PyTorch中GPU OOM问题原因探究
问题场景
以下PyTorch代码片段是损失计算的一部分,运行多个epoch后会触发GPU内存不足(OOM):
eigvecs = torch.randn(n, b, dtype=eigval_approximations.dtype, device=eigval_approximations.device) eigvecs /= torch.linalg.norm(eigvecs, dim=0) for _ in range(iterations): eigvecs = torch.linalg.solve(mats - eigval_approximations * identity, eigvecs.T).T eigvecs /= torch.linalg.norm(eigvecs, dim=0)
其中eigvecs在每次训练/验证步骤都会重新初始化,尺寸固定且作用域短暂。将代码中的原地除法(/=)改为非原地操作(eigvecs = eigvecs / torch.linalg.norm(...))后,内存泄漏问题完全消失。
疑问:为何原地操作会引发内存泄漏?这是Python的预期行为,还是PyTorch的实现细节或bug?
原因分析
这并非Python的预期行为,而是PyTorch自动微分机制与内存管理的实现细节导致的,不属于PyTorch的bug,核心原因如下:
计算图梯度跟踪的引用残留:PyTorch默认会为张量开启梯度跟踪(除非显式设置
requires_grad=False),当使用原地操作(如/=)时,会直接修改张量本身的数据,但计算图为了支持梯度反向传播,会保留该张量的历史版本引用。这些旧版本张量无法被Python垃圾回收器及时标记回收,也无法被CUDA内存分配器释放,随着epoch累积,GPU内存占用持续上升最终触发OOM。非原地操作的内存回收优势:改用非原地操作(
eigvecs = eigvecs / torch.linalg.norm(...))时,会创建一个全新的张量对象,原来的旧张量因为不再被任何引用(包括计算图)持有,会被Python垃圾回收器快速标记,进而被CUDA内存分配器回收,不会形成内存堆积。原地操作破坏张量不可变性设计:PyTorch的自动微分依赖张量的不可变性来维护计算图的正确性,原地操作会修改张量的
_version属性(用于跟踪张量修改次数),迫使PyTorch保留更多的历史状态信息以确保梯度计算的准确性,这进一步加剧了内存占用。
内容的提问来源于stack exchange,提问作者Moon

