PyTorch中non_blocking从CUDA到CPU赋值结果异常问题
GPU到CPU张量切片赋值(
non_blocking=True)结果异常的解决方案 问题原因分析
你遇到的核心问题是**non_blocking=True在GPU到CPU异步传输时的时机错误**:
- 启用
non_blocking=True后,b.to(cpu, ...)会立即返回一个CPU张量,但数据并未完成从GPU到CPU的拷贝,此时的CPU张量内存是未初始化或未更新的状态。 - 原代码中
torch.cuda.synchronize()在赋值操作之后执行,无法修正已经完成的赋值——赋值时读取的是未完成传输的内存,导致结果滞后、出现0值或随机垃圾值(比如负数)。
已知copy_方法速度慢、直接覆盖变量会增加内存开销,下面提供两种兼顾性能与正确性的解决方案:
解决方案1:Pin Memory + 先同步后赋值
给CPU张量启用pin_memory(锁页内存),配合异步传输+提前同步,既能保留non_blocking的传输效率,又能保证数据正确性:
import torch gpu = torch.device('cuda') cpu = torch.device('cpu') # 初始化CPU张量时启用pin_memory,支持GPU到CPU的DMA异步传输 a = torch.rand((13223,134,4), dtype=torch.float32, device=cpu, pin_memory=True) b = torch.rand((13223,134,4), dtype=torch.bfloat16, device=gpu) for i in range(3): b.mul_(0.5) # 执行异步传输到临时CPU张量 temp = b.to(device=cpu, memory_format=torch.preserve_format, dtype=torch.bfloat16, non_blocking=True) # 等待传输完成后,再执行切片赋值 torch.cuda.synchronize() a[:] = temp print(b[0,0], a[0,0])
效果
运行后输出的GPU张量与CPU张量数值完全对齐:
tensor([0.0942, 0.1621, 0.2041, 0.1543], device='cuda:0', dtype=torch.bfloat16) tensor([0.0942, 0.1621, 0.2041, 0.1543]) tensor([0.0471, 0.0811, 0.1021, 0.0771], device='cuda:0', dtype=torch.bfloat16) tensor([0.0471, 0.0811, 0.1021, 0.0771]) tensor([0.0236, 0.0405, 0.0510, 0.0386], device='cuda:0', dtype=torch.bfloat16) tensor([0.0236, 0.0405, 0.0510, 0.0386])
解决方案2:显式CUDA流同步(适合多任务场景)
如果你的程序同时有其他GPU计算任务,全局同步torch.cuda.synchronize()会阻塞所有GPU操作,此时可以用单独的CUDA流做数据传输,仅同步该流,避免影响其他任务:
import torch gpu = torch.device('cuda') cpu = torch.device('cpu') a = torch.rand((13223,134,4), dtype=torch.float32, device=cpu, pin_memory=True) b = torch.rand((13223,134,4), dtype=torch.bfloat16, device=gpu) # 创建独立的CUDA流用于数据传输 transfer_stream = torch.cuda.Stream() for i in range(3): b.mul_(0.5) # 在传输流中执行异步传输 with torch.cuda.stream(transfer_stream): temp = b.to(device=cpu, memory_format=torch.preserve_format, dtype=torch.bfloat16, non_blocking=True) # 仅同步传输流,不阻塞其他GPU计算操作 transfer_stream.synchronize() a[:] = temp print(b[0,0], a[0,0])
关键优化点
- Pin Memory:锁页内存允许GPU直接通过DMA访问CPU内存,避免了内存页交换的开销,大幅提升异步传输速度。
- 同步时机:必须在赋值操作之前完成传输同步,确保CPU张量的数据已经完全更新。
- 显式流控制:在多GPU任务场景下,显式流能最小化同步对其他计算任务的影响,提升整体并行效率。
内容的提问来源于stack exchange,提问作者Anonymous
相关产品推荐
相关产品推荐

