You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何使用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()

关键改进点说明

  1. 移除Python线程:CUDA流本身就是异步执行的,不需要额外线程来实现并行传输和计算,线程反而会破坏流的上下文绑定。
  2. 正确建立流依赖:
    • s1.wait_stream(default_stream):确保只有当前Tensor计算完成后,才开始回传CPU,避免读取未完成计算的数据。
    • default_stream.wait_stream(s2):确保下一个Tensor传输完成后,才开始GPU计算,避免使用未加载完成的数据。
  3. 正确的Tensor操作:直接使用to(device, non_blocking=True)进行异步传输,避免操作data属性导致的Tensor引用丢失。
  4. 合理的record_stream使用:在GPU Tensor传输完成后调用record_stream,告诉PyTorch该Tensor还被对应流使用,防止显存被提前回收;回传CPU后同样标记,确保GPU端的Tensor副本可以被正确释放。
  5. 简化循环逻辑:采用索引循环实现流水线,避免重复处理同一个Tensor,减少不必要的显存占用。

实际场景OOM的原因

在你的实际应用中,因为有数千个Tensor,且流的依赖关系失效,导致GPU同时加载了大量未被处理的Tensor,而PyTorch的显存回收机制无法识别这些可释放的内存(因为record_stream使用错误,流的同步失效),最终导致显存持续增长直至溢出。修复流的依赖关系和Tensor操作逻辑后,显存会随着流水线的推进及时回收,避免OOM问题。

内容的提问来源于stack exchange,提问作者Anonymous

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.18 12:07:05