PyTorch训练TranX文本转代码模型时loss.backward阶段CPU/GPU内存泄漏
内存泄漏排查方向及解决方案
1. 循环引用导致的内存无法回收
你定义的Example类、AST生成对象存在大量循环引用场景:
- AST节点之间互相持有父/子节点引用,
tgt_ast属性挂载到Example后形成闭环引用 Example.meta字段存储了全量原始行数据、slot映射等冗余信息,和上层对象形成引用链
Python GC对绑定了PyTorch张量的循环引用回收效率极低,甚至会出现永久无法释放的情况。
解决方案:- 数据预处理阶段仅保留训练必需的输入输出字段,返回
Example前显式del所有中间AST对象、原始行数据,断开引用链 - 不要将不需要参与梯度计算的非原生Python对象绑定到返回的样本结构中
2. model.score方法隐性缓存中间张量
原TranX仓库的score方法大概率将中间计算张量挂载到了model实例属性上做缓存,你仅清理了方法内的局部变量,但model实例持有的张量引用一直存在,每轮迭代都会新增一批缓存对象。
解决方案:
- 检查
model实例的__dict__属性,确认每次调用score后是否新增了非参数类的张量属性 - 每次调用
score后手动清理所有临时缓存属性,不要依赖方法内的局部变量del操作
3. PyTorch版本已知Bug
PyTorch 1.9.1存在多个autograd反向传播阶段的内存泄漏已知问题,尤其在处理自定义非结构化计算图、变长序列输入时会触发显存/内存泄漏。
解决方案:
- 升级PyTorch版本至1.12.x及以上稳定版,该版本修复了大量反向传播相关的内存泄漏问题
4. 全量验证导致的内存溢出
你验证阶段直接将全量验证集样本传入evaluate函数,若该函数内部缓存了所有样本的中间预测结果、梯度张量,会一次性占用大量内存/显存。
解决方案:
- 验证阶段改用分批推理,每批推理完成后立即清理所有中间变量,不要累积全量验证集的计算结果
5. 冗余数据常驻内存
你的数据集实现中将全量原始文本行一次性读入内存存储,__getitem__返回的Example携带大量训练用不到的冗余字段,迭代数万次后冗余数据累积直接占满32GB系统内存,触发CPU OOM。
解决方案:
- 预处理阶段提前把所有样本需要的字段提取出来存储,不要在训练时实时解析全量json行、生成冗余AST对象
- 若数据集过大,改用流式读取方式加载数据,不要一次性全量读入内存
内容的提问来源于stack exchange,提问作者abtExp
相关产品推荐
相关产品推荐

