PyTorch使用.cuda()时GPU显存占用过高的原因及优化方法咨询

调用
.cuda()时显存过高的原因 - CUDA上下文初始化固定开销:第一次调用
.cuda()时,PyTorch会自动加载CUDA驱动、cuDNN等核心依赖库,同时默认预分配一块显存作为缓存池减少后续申请开销,这部分固定占用通常在500MB~1.2GB区间,对于显存只有6GB的GTX 1660来说占比会非常明显。 - 模型参数直接占用:迁移到GPU的模型权重、偏置全部会加载到显存,例如参数量过亿的大模型仅权重就会占用数GB显存。
- 中间张量隐式迁移:如果调用
.cuda()前CPU侧有绑定计算逻辑的张量,迁移时会同步将后续计算需要的中间张量一次性加载到显存,造成瞬时占用飙升。 - 残留显存未释放:之前异常退出的PyTorch进程可能没有释放占用的显存,新进程启动后占用会叠加,导致显存占用远超预期。
降低显存占用的可行方法
前置优化
- 清理残留显存:执行
nvidia-smi命令查看显存占用的进程ID,用kill -9 <进程ID>终止僵尸进程释放显存,也可以直接重启GPU驱动清空所有残留。 - 关闭显存预分配:运行代码前设置环境变量
export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,关闭默认的大块显存预分配策略,仅在需要时动态申请显存,初始占用可以降低30%左右。 - 低精度加载模型:默认FP32精度的模型改为FP16或INT8精度加载,显存占用直接降低50%~75%,示例代码如下:
# 转为FP16精度 model = model.half().cuda() # 动态量化为INT8精度,适配推理场景 model = torch.ao.quantization.quantize_dynamic(model, {torch.nn.Linear}, dtype=torch.qint8).cuda()
运行时优化
- 避免冗余张量迁移:不需要参与GPU计算的张量不要调用
.cuda(),用完的张量及时用del 张量名删除,随后调用torch.cuda.empty_cache()清空闲置缓存。 - 推理场景开启无梯度上下文:所有推理代码外层套
with torch.no_grad():,关闭梯度记录功能,无需保存反向传播所需的中间张量,显存占用可降低50%以上。 - 训练场景调优:
- 降低批量大小(batch size),单步计算的样本数减少后,对应的中间张量占用会线性下降
- 用梯度累积替代大batch,每N步计算后再统一更新梯度,在不改变训练效果的前提下,显存仅需满足单步小batch的占用需求
- 开启梯度检查点,用
torch.utils.checkpoint.checkpoint包装模型层,通过重计算代替中间张量存储,最高可降低70%的训练显存占用
内容的提问来源于stack exchange,提问作者y chen
相关产品推荐
相关产品推荐

