如何删除PyTorch计算图中的张量?目标检测循环预测张量残留如何避免
PyTorch for循环预测时计算图残留问题解决方法
产生该问题的核心原因是PyTorch默认运行模式会跟踪所有张量的计算过程,存储计算图中间张量用于后续反向传播,即便推理场景下没有调用反向传播接口,这些关联张量也不会被自动回收,可通过以下方法修复:
- 模型实例化后切换为评估模式,关闭Dropout、BatchNorm等训练专属逻辑,同时减少不必要的中间变量存储:
model = detection.fasterrcnn_resnet50_fpn(pretrained=True, progress=True,pretrained_backbone=True).to(DEVICE) model.eval() # 新增评估模式代码
- 预测逻辑外层包裹推理上下文,彻底禁用计算图跟踪能力,完全不存储和计算梯度相关的中间变量,推荐使用性能更优的
torch.inference_mode(),也可使用常用的torch.no_grad() - 单次预测完成后可手动删除不再使用的临时张量,若使用GPU设备可主动调用显存清空接口,避免显存碎片残留。
修改后的完整代码如下:
model = detection.fasterrcnn_resnet50_fpn(pretrained=True, progress=True,pretrained_backbone=True).to(DEVICE) model.eval() # 切换评估模式 with torch.inference_mode(): # 新增推理上下文 for i in tqdm(range(train.shape[0])): image = cv2.imread(train_img_paths[i]) image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) image = image.transpose((2, 0, 1)) image = image / 255.0 image = np.expand_dims(image, axis=0) image = torch.FloatTensor(image) image = image.to(DEVICE) predictions = model(image)[0] # 后续处理完predictions后可手动清理变量 del image, predictions # 若使用GPU可加如下代码清空缓存 # torch.cuda.empty_cache()
如果你需要在TensorFlow下解决同类问题,只要在预测逻辑外层套tf.stop_gradient(),或者使用tf.function(jit_compile=True)修饰预测函数,同时调用模型时传入training=False参数即可。
内容的提问来源于stack exchange,提问作者Olli
相关产品推荐
相关产品推荐

