双RTX3090 GPU训练时反向传播CUDA内存不足问题排查与优化
问题背景
使用两块NVIDIA RTX 3090(单卡24GB显存)训练模型,将约6GB的大型稀疏矩阵加载至cuda:1,模型部署在cuda:0。即便batch size设为1,loss.backward()阶段仍触发CUDA内存不足错误。将稀疏矩阵移至CPU后代码可正常运行,但数据传输开销导致速度变慢;此时观察到loss.backward()阶段RAM占用增加约9GB(总CPU RAM为40GB)。
简化训练代码
def train_model(model, data_loader, num_epochs, sparse_matrix_path, save_path): device_net = torch.device('cuda:0') device_matrix = torch.device('cuda:1') model = model.to(device_net) # Load model and sparse matrix model.train() sparse_matrix = torch.load(sparse_matrix_path, weights_only=True).to_sparse().to(device_matrix) sparse_matrix.requires_grad = False optimizer = torch.optim.Adam(model.parameters(), lr=1e-4) for epoch in range(num_epochs): for step, input_data in enumerate(data_loader): input_data = input_data.to(device_net).unsqueeze(1) optimizer.zero_grad() # Model forward pass output, _ = model(input_data) # Sparse matrix operation batch_size = output.size(0) flat_output = output.view(batch_size, -1).to(device_matrix) sparse_output = torch.sparse.mm(sparse_matrix, flat_output.T) # Compute loss transformed_output = sparse_output.T.view(batch_size, 1, 700, 384) loss = F.mse_loss(transformed_output.to(device_net), input_data) loss.backward() # Error occurs here optimizer.step()
关键观察
- 6GB稀疏矩阵部署在
cuda:1 - CPU运行时,
loss.backward()使RAM占用增加约9GB - GPU运行时,
loss.backward()阶段显存占用骤增,触发内存不足错误
核心问题
- 为何存放稀疏矩阵的GPU在
loss.backward()阶段显存占用大幅增加? - 如何优化该GPU部署方案以避免内存不足?有哪些PyTorch功能或策略可降低反向传播阶段的GPU内存占用?
问题解答
1. 显存占用骤增的原因
虽然设置了sparse_matrix.requires_grad = False,但反向传播过程中,稀疏矩阵与flat_output的矩阵乘法操作(torch.sparse.mm)会生成中间张量的梯度计算图。即便稀疏矩阵本身不需要梯度,flat_output是来自cuda:0的模型输出,它需要梯度回传,而矩阵乘法的中间结果会在cuda:1上保留梯度相关的缓存。
另外,sparse_output及其后续的转置、变形操作都会在cuda:1上生成临时张量,反向传播时这些张量的梯度信息会占用额外显存。从CPU运行时RAM增加9GB也能印证:反向传播需要存储大量中间结果的梯度,对应到GPU上就是显存占用暴涨。
2. 优化方案与PyTorch策略
(1)隔离稀疏矩阵操作的计算图
在稀疏矩阵操作部分使用torch.no_grad()包裹,避免存储这部分操作的中间梯度缓存(不影响模型参数的梯度回传):
# Sparse matrix operation batch_size = output.size(0) flat_output = output.view(batch_size, -1).to(device_matrix) with torch.no_grad(): sparse_output = torch.sparse.mm(sparse_matrix, flat_output.T)
(2)使用梯度检查点减少显存占用
通过torch.utils.checkpoint在反向传播时重新计算中间结果,而非存储它们,大幅降低显存开销:
from torch.utils.checkpoint import checkpoint # 封装稀疏矩阵操作 def sparse_op(sparse_mat, flat_out): return torch.sparse.mm(sparse_mat, flat_out.T) # 训练循环中替换原操作 flat_output = output.view(batch_size, -1).to(device_matrix) sparse_output = checkpoint(sparse_op, sparse_matrix, flat_output)
(3)优化张量传输与内存复用
- 提前预分配
flat_output的显存空间,避免每次迭代重新分配:循环外创建合适大小的张量,循环内直接复用。 - 对
sparse_output的转置、变形操作优先使用原地操作(如.transpose(0,1).contiguous()替代.T.contiguous()),减少临时张量的创建。
(4)严格控制梯度范围
确保sparse_matrix的requires_grad始终为False,避免PyTorch为其分配梯度空间;反向传播时设置retain_graph=False(默认值),及时释放无用的计算图:
loss.backward(retain_graph=False)
(5)精准监控显存使用
在关键节点(forward后、backward前后)打印显存占用,定位具体的内存开销来源:
# 打印cuda:1的显存占用 print(f"Before backward: {torch.cuda.memory_allocated(device_matrix)/1e9:.2f} GB") loss.backward() print(f"After backward: {torch.cuda.memory_allocated(device_matrix)/1e9:.2f} GB")
内容的提问来源于stack exchange,提问作者Bahareh

