如何使用PyTorch CUDA并行流避免OOM?附示例代码
CUDA多流异步传输导致显存溢出(OOM)的问题分析与修复
核心问题分析
你的代码存在几个关键错误,导致实际场景中显存无法正确回收,最终触发OOM:
- Python线程与CUDA流上下文不兼容:CUDA流的上下文是线程局部的,你用
Thread启动的函数中,torch.cuda.stream(s1)无法正确绑定到主线程的流对象,导致流的同步逻辑完全失效。异步传输和计算之间没有正确的依赖关系,GPU上会堆积大量未处理的Tensor副本。 - 流同步逻辑混乱:
s1.wait_stream(default_stream)和s2.wait_stream(s1)的依赖关系完全错误,s1.synchronize()会阻塞整个进程,彻底失去异步流水线的意义。计算流(默认流)和传输流之间没有建立正确的等待关系,导致GPU同时加载过多数据。 - Tensor操作与record_stream使用错误:
p.data = p.data.to(...)会创建新的Tensor对象,原GPU Tensor的引用被丢失,PyTorch的显存回收机制无法正确识别可释放的内存。同时,CPU Tensor无法调用record_stream,你在GPU传输后的record_stream时机也不对,应该在GPU Tensor上操作,告诉PyTorch该Tensor还被指定流使用。 - 冗余循环导致重复处理:外层无限循环和内层
j的循环范围会重复处理同一个Tensor多次,加剧显存堆积。
修复后的代码示例
以下是修正后的实现,去掉了不必要的线程,正确建立流依赖,实现GPU计算、上块回传CPU、下块传GPU的流水线并行:
import torch from time import perf_counter cpu = torch.device('cpu') gpu = torch.device('cuda') # 初始化10个CPU Tensor tensor_count = 10 tensors = [torch.rand(100_000_000, device=cpu) for _ in range(tensor_count)] # 定义CUDA流:s1负责回传CPU,s2负责传输到GPU s1 = torch.cuda.Stream(device=gpu) s2 = torch.cuda.Stream(device=gpu) default_stream = torch.cuda.default_stream(gpu) # 预加载第一个Tensor到GPU current_tensor = tensors[0].to(gpu, non_blocking=True) current_idx = 0 # 流水线循环处理 for _ in range(100): # 替换为实际循环次数 time_start = perf_counter() # 1. 异步将上一个处理完的Tensor传回CPU(如果不是第一个) if current_idx > 0: prev_tensor = tensors[current_idx - 1] with torch.cuda.stream(s1): # 确保s1等待默认流完成计算,才能读取结果 s1.wait_stream(default_stream) # 异步回传CPU prev_tensor.data = prev_tensor.to(cpu, non_blocking=True) # 标记该Tensor不再被s1使用,允许显存回收 prev_tensor.record_stream(s1) # 2. 异步加载下一个Tensor到GPU(如果不是最后一个) next_idx = (current_idx + 1) % tensor_count with torch.cuda.stream(s2): next_tensor = tensors[next_idx] # 异步传输到GPU next_tensor.data = next_tensor.to(gpu, non_blocking=True) # 标记该Tensor被s2使用,防止提前释放 next_tensor.record_stream(s2) # 3. 在默认流处理当前GPU Tensor(等待s2完成传输) default_stream.wait_stream(s2) with torch.cuda.stream(default_stream): for _ in range(10): current_tensor.mul_(1.01) current_tensor.add_(1.01) current_tensor.pow_(0.5) # 更新当前Tensor索引 current_idx = next_idx current_tensor = tensors[current_idx] # 统计耗时 lapsed = perf_counter() - time_start print(f"\rspeed: {lapsed:.5f} sec/it", end="\r") # 最后同步所有流,确保所有操作完成 torch.cuda.synchronize()
关键改进点说明
- 移除Python线程:CUDA流本身就是异步执行的,不需要额外线程来实现并行传输和计算,线程反而会破坏流的上下文绑定。
- 正确建立流依赖:
s1.wait_stream(default_stream):确保只有当前Tensor计算完成后,才开始回传CPU,避免读取未完成计算的数据。default_stream.wait_stream(s2):确保下一个Tensor传输完成后,才开始GPU计算,避免使用未加载完成的数据。
- 正确的Tensor操作:直接使用
to(device, non_blocking=True)进行异步传输,避免操作data属性导致的Tensor引用丢失。 - 合理的record_stream使用:在GPU Tensor传输完成后调用
record_stream,告诉PyTorch该Tensor还被对应流使用,防止显存被提前回收;回传CPU后同样标记,确保GPU端的Tensor副本可以被正确释放。 - 简化循环逻辑:采用索引循环实现流水线,避免重复处理同一个Tensor,减少不必要的显存占用。
实际场景OOM的原因
在你的实际应用中,因为有数千个Tensor,且流的依赖关系失效,导致GPU同时加载了大量未被处理的Tensor,而PyTorch的显存回收机制无法识别这些可释放的内存(因为record_stream使用错误,流的同步失效),最终导致显存持续增长直至溢出。修复流的依赖关系和Tensor操作逻辑后,显存会随着流水线的推进及时回收,避免OOM问题。
内容的提问来源于stack exchange,提问作者Anonymous
相关产品推荐
相关产品推荐

