☰
多流推理必读:record_stream 和 wait_event 的职责与正确用法
2026/10/6 5:57:27 网站建设 项目流程

一次线上推理服务的延迟优化,把原本跑在默认流上的处理流程拆成了三路 CUDA Stream:一路负责 CPU 侧加载和预处理,一路负责 H2D 拷贝,一路专门跑模型推理。改完一测,第一个版本就翻车了——不是立刻崩,而是输出里时不时出现 NaN 和错位帧,偶尔还整卡报 illegal memory access。排查了两天,断点最终落在tensor.record_stream()和stream.wait_event()这两个 API 的职责区别上。

这篇就完整记录这次掉坑过程:为什么只记住record_stream保命还不够,为什么多流场景下wait_event才是真正把顺序钉死的那颗钉子,以及现在我在工程里固定使用的多流模板。适合正在做多流推理、数据预加载、或任何想把张量交给主流之外的流去消费的同学参考。

1. 先说结论:两个 API,管的是两件完全不同的事

1.1 record_stream 到底在做什么

官方文档原话:Ensures that memory from the tensor is not reused until all current work on the given stream is complete。意思是:保证张量背后的显存,在给定流上所有当前工作完成之前,不会被缓存分配器提前回收或复用。

要理解这条,得先了解 PyTorch 的 CUDA Caching Allocator。默认情况下,为了砍掉反复cudaMalloc的开销,分配器会从一个大块 segment 里给张量划出小块 block,张量析构时内存块回到缓存池,并不真正还给驱动。问题就在“回池”的时机:PyTorch 默认认为,张量只会在创建它的流(也就是分配那一刻上下文里的 current stream)上使用。如果这个张量被拿到另一个流上去消费,消费还没完成,Python 端的引用先归零了,那么分配器极有可能把内存块放回池子,立刻被下一个torch.randn、torch.zeros申请接走。此时原流上那些正在读这块显存的 kernel 还没跑完,数据已经被覆盖,数值错误几乎无法避免。

record_stream(usage_stream)就是把这个“另一个流”的编号告诉分配器:这块内存这个流上还在用,先别急着回收。它本质上是在 usage_stream 上记录了一个内部的释放事件,内存块要等所有这些事件都完成后,才真正允许复用。

1.2 wait_event / wait_stream 在做另一件事

stream.wait_event(event)/stream.wait_stream(other_stream)建立的是跨流执行的 happens-before 关系:目标流上排队的所有后续工作,都必须等到源流某个事件完成之后才能执行。这是 CUDA 事件机制最基础也最核心的用法。

我常用的事件同步写法:

ev = torch.cuda.Event() # 在源流上记录:代表"源流此前的工作到此为止" with torch.cuda.stream(s1): ev.record(s1) # 在目标流上等待:让 s2 之后的执行排在 ev 之后 with torch.cuda.stream(s2): s2.wait_event(ev)

有时也可以直接s2.wait_stream(s1),语义是“把 s1 当前所有未完成工作用一个事件包装,s2 等待”,省去手动建事件。实际工程里两者我都用:如果只想等待源流上一个阶段的工作,就用wait_event+ 手工插桩;如果就是等源流全部干完,wait_stream更省事。

1.3 为什么很多人只写了 record_stream 就翻车

因为record_stream对执行顺序没有任何约束力,它只是分配器层面的一条“暂缓回收”指令。你可以脑补这样一个糟糕的时间线:

  • s1 上的 kernel 刚把数据写入张量内存,还没写完;
  • Python 引用计数归零,块被record_stream保护,分配器没有立刻回收;
  • s2 上的消费 kernel 排入队列,但 s2 没有等待 s1,硬件调度直接让消费 kernel 先跑或者乱序执行;
  • 消费 kernel 读到一半甚至全空的内存,结果自然是错的。

record_stream保护的是“内存不被抢走”,wait_event保证的是“顺序不乱”。内存活下来了,但没有顺序保证,数据竞争照样在。这就像你把一个仓库划给两个团队用:record_stream是跟物业说“东西还占着仓库,到期前别把钥匙给别人”,wait_event是跟施工方说“必须等上一家搬完再进门”。两个都少不得。

2. 我踩坑的原始场景:一个典型的多流推理改造

2.1 改造前的单流代码长什么样

我做这个推理服务,输入是一批视频帧,每帧要完成 resize、归一化、H2D 拷贝、模型推理、D2H 拷贝,然后再拼帧输出。最早的实现就是全部在默认流上串行跑,吞吐很容易卡在 CPU 预处理和拷贝等待上。后来把逻辑拆成 load_stream(CPU 预处理 + 拷贝)、compute_stream(模型)、copy_back_stream(结果回传),三个流形成一条简单的三级流水线。

起初的伪代码长这样:

load_stream = torch.cuda.Stream() compute_stream = torch.cuda.Stream() copy_back_stream = torch.cuda.Stream() for batch in frames_batches: # CPU 预处理 + H2D 拷贝 with torch.cuda.stream(load_stream): inputs_gpu = preprocess(batch).to(device, non_blocking=True) # 模型推理 with torch.cuda.stream(compute_stream): outputs = model(inputs_gpu) # 结果回传 with torch.cuda.stream(copy_back_stream): results = outputs.to('cpu', non_blocking=True)

当时我天真地以为,进入不同 stream 的上下文,流之间自然就“同步”了。实际上完全不是:with torch.cuda.stream(s)只是切换当前流,“当前流之前的工作”不会自动等别的流。三个流之间的执行顺序是完全独立的,要靠显式同步去钉。

2.2 引入多流之后出现的诡异症状

改造后第一次完整跑,整个链路的表现是:不崩还好,一崩就崩得很难看。症状大概分三类:

  • 输出里偶发 NaN 和明显错位的帧,batch 越大概率越高。这种最头疼,因为你不知道是数据问题、模型问题还是流问题。
  • 偶尔 CUDA 直接抛illegal memory access,或者 PyTorch 报CUDA error: device-side assert triggered。一开始我还以为是自己的代码里哪个 index 越界了。
  • 还有一个隐蔽问题:显存占用莫名其妙地涨。因为non_blocking拷贝没等待,事件和临时张量互相纠缠,分配器里的块迟迟释放不掉。

这三类症状在很多多流优化项目里都很典型:数据竞争、生命周期错乱、分配器恐慌。总结起来就一句话:没有把跨流依赖关系钉死,硬件就会给你自由发挥的机会。

2.3 复现步骤与最小化测试

遇到这种问题,我一般会先写一个最小化复现脚本,把代码压缩到几十行,排除模型、数据加载器的干扰。核心逻辑就是:s1 上创建张量,s2 上消费,循环跑很多次,看数值是否稳定。

import torch s1 = torch.cuda.Stream() s2 = torch.cuda.Stream() for i in range(1000): with torch.cuda.stream(s1): a = torch.randn(4096, 4096, device='cuda') with torch.cuda.stream(s2): b = a * 2.0 if not torch.isfinite(b).all(): print(f"i={i}: got non-finite in s2 result") break

这段代码我在不同设备上都能复现出错误:s2 上的 kernel 可能会先于 s1 执行,b计算出来的值有概率是 NaN。把等待逻辑补上之后,结果稳定。这个最小化脚本后来也成了我给团队同学演示“多流为什么必须同步”的标准案例。

3. 排查过程:从怀疑数据到锁定流同步

3.1 第一层排查:去掉多流是否正常

我排查这种问题时,第一步永远是做“二分对照”:把所有with torch.cuda.stream上下文全部去掉,回归到默认流串行版本。如果串行版本稳定,多流版本不稳定,那方向就锁定在流同步、内存生命周期这两者上,而不是模型和数据本身。

第二步是把三个流一分为二地启用,比如只让 load_stream 和 compute_stream 协作,看问题是否保留。这一步能缩小到具体是哪条流关系没理顺。我在实际项目里就是这样一步步从“怀疑模型代码”走到“怀疑内存管理”的。

3.2 第二层排查:事件状态与 stream 顺序

当我把问题锁定在流同步后,开始用事件做探针,检查流的执行时机。做法:在源流记录事件,用事件查询看状态变化,然后在目标流里看结果是否稳定。

ev = torch.cuda.Event() with torch.cuda.stream(s1): a = torch.randn(4096, 4096, device='cuda') ev.record(s1) print("event completed:", ev.query()) # False 说明还没完成 with torch.cuda.stream(s2): b = a * 2.0 # 观察 b 是否稳定

关键观察:如果不加s2.wait_event(ev),ev是否完成和 s2 上b的计算顺序没有必然关系。CUDA 驱动对独立流的调度是自由的,谁先谁后由驱动决定。这时候我意识到问题不是“数据没拷完”这种初级错误,而是“两个流的执行顺序压根没被约束”。

3.3 第三层排查:内存分配器在“捣乱”

接下来我要验证“内存被提前复用”是否也在出力。方法很简单:给 s1 创建的张量调record_stream(s2)和不调,分别看结果。调了之后,NaN 概率还是会存在(因为没 wait_event),但显存异常增长和 illegal memory access 明显变少。这说明record_stream对应的回收保护起了作用,但没有解决顺序问题。

另外一个更直观的实验:让 s2 的消费操作人为慢一点,比如在 s2 里加一个torch.cuda.synchronize()或短暂等待,错误就消失了。这基本验证了“流间缺少同步”这个根因。

这三次排查,最后落到一个结论:所有问题都指向同一个缺失——目标流上没有等待源流。

4. 根因详解:缓存分配器、事件与执行顺序的真实关系

4.1 PyTorch CUDA Caching Allocator 的工作机制

要彻底理解这个问题,得深入一点 PyTorch 的显存缓存分配器。

PyTorch 从驱动拿到的显存是一次性、大块的 segment,内部再切成小块 block。分配器维护每个 block 的状态:use_count、stream 关联等。张量析构时,如果 block 上没有未完成的 stream 使用记录,就直接回到 free 列表,等待下一次分配复用。

问题在于,一个 block 被创建时记录的“默认使用流”,就是创建它的上下文里的 current stream。当张量被传到其他流上时,如果不告诉分配器,分配器就不知道其他流也在排队等这块内存,于是放心地把 block 回收。

record_stream正是为此设计的:它会把 block 追加到 stream_uses 集合,并在张量析构时把 block 放入延迟释放队列,同时在该 stream 上记录一个事件。只有这些事件都在驱动层完成之后,block 才会被真正送回 free 列表。

需要注意一个实现细节:record_stream只是延长生命周期,不会给目标流插入任何 wait。这也是我踩坑的核心——我一开始确实写了record_stream,但漏了 wait_event,以为内存保住了就万事大吉。结果内存没被复用,但 s2 上的 kernel 仍然可能早于 s1 执行,写读竞争照样产生错误数据。

4.2 record_stream 的内部实现与局限

复盘时我专门去翻了 PyTorch 源码。Python 层的Tensor.record_stream(stream)最终会走到CUDACachingAllocator::recordStream。实现要点:

  • 根据张量指针找到缓存块;
  • 如果该块当前正被另一个流使用,就把新流加进 stream_uses 集合;
  • 当块释放时,会在这个集合里的每个流上记录一个事件,专门用于延迟回收。

所以它管不到“计算顺序”:record_stream(usage_stream)里的 usage_stream 只是“告诉缓存分配器哪个流在用”,它并不负责让 usage_stream 的队列等待源流完成。真正需要等待,必须靠wait_event/wait_stream另行声明。

我复盘时也重新确认了 PyTorch 文档里Tensor.record_stream的说明,文档从没说过它会造成跨流等待。所以这个坑更多是“直觉上觉得它有用”导致的:看着名字像是要“记录流”,以为它会负责流同步,其实它只管生命周期。

4.3 wait_event 补上的那一环

wait_event的本质是往目标流的命令流里插入一个“等待点”:源流上对应事件之前的所有工作,都要先完成,目标流上等待点之后的所有 kernel 才能开始执行。

如果把两条流比作两条生产线:

  • 源流向事件里投递一个“这批货打包完成”的信号;
  • 目标流在事件上做一次“等信号”的停靠;
  • 信号不到,目标流的后续工序不启动。

这就是 happens-before 的具象化。

需要补充一个细节:wait_event的等待点和record_stream的生命周期注册是“各管各的”。即使你写了wait_event,也不能代替record_stream。因为等待点只保证执行顺序,不保证“这块内存的回收时间”。特别是张量的 Python 引用可能在目标流执行完成前就被释放,缓存分配器仍然可能把内存复用给别的流。所以正确姿势是两者都写,职责互补。

4.4 完整正确的模板代码

下面是我现在工程里固定使用的多流协作模板,以数据加载 + 模型推理为例:

# 创建三路流和事件 load_stream = torch.cuda.Stream() compute_stream = torch.cuda.Stream() copy_back_stream = torch.cuda.Stream() h2d_done = torch.cuda.Event() compute_done = torch.cuda.Event() for i in range(total_batches): # 1) CPU 预处理 + H2D 拷贝(异步) with torch.cuda.stream(load_stream): inputs_gpu = preprocess(next(iter(dataloader))).to(device, non_blocking=True) h2d_done.record(load_stream) # 2) 计算流必须等拷贝流完成 with torch.cuda.stream(compute_stream): compute_stream.wait_event(h2d_done) outputs = model(inputs_gpu) compute_done.record(compute_stream) inputs_gpu.record_stream(compute_stream) # 3) 回传流必须等计算流完成 with torch.cuda.stream(copy_back_stream): copy_back_stream.wait_event(compute_done) results = outputs.to('cpu', non_blocking=True) outputs.record_stream(copy_back_stream)

模板的关键点:

  • 每个消费流在进入业务逻辑前,先wait_event上游事件;
  • 每个在非创建流上使用的张量,紧随使用点调用record_stream;
  • 事件的记录位置,在源流完成所有生产工作之后;
  • record_stream尽量早、尽量紧贴使用点,避免张量引用被释放很久后才发现忘了保护。

这套模板我在好几个项目里直接用,没再踩过同类坑。

5. 工程中的常见场景与正确姿势

5.1 数据预加载 + 计算

数据预加载是最容易踩坑的区域,因为DataLoader的 worker 和主进程之间的张量搬运天然异步,很多人还在主线程里手动用 pinned memory + non_blocking。

常见的做法是:

with torch.cuda.stream(prefetch_stream): next_batch = batch.to(device, non_blocking=True) # 计算时切换回主流,但主流必须等待 prefetch 完成 torch.cuda.current_stream().wait_stream(prefetch_stream) out = model(next_batch)

这里最容易漏的是next_batch.record_stream(torch.cuda.current_stream())。因为next_batch是在 prefetch_stream 上创建的,但主流上的 kernel 也在消费它。如果next_batch在这个迭代结束时被覆盖,内存可能被提前回收。我建议在模型消费完的下一行立刻补record_stream,别拖到函数末尾再补,那样一旦中间有早返回分支就遗漏了。

5.2 多流并行推理 + 汇总

另一种高频场景是把一个 batch 拆成多个 chunk,每个 chunk 放到独立 stream 上并行推理,最后在主流汇总。

streams = [torch.cuda.Stream() for _ in range(num_streams)] events = [torch.cuda.Event() for _ in range(num_streams)] results = [None] * num_streams chunks = torch.chunk(big_batch, num_streams) for idx, chunk in enumerate(chunks): with torch.cuda.stream(streams[idx]): # 子流要等主流的 chunk 切分完成 streams[idx].wait_stream(torch.cuda.current_stream()) results[idx] = model(chunk) chunk.record_stream(streams[idx]) events[idx].record(streams[idx]) # 主流汇总前,等待所有分支流 for idx in range(num_streams): torch.cuda.current_stream().wait_event(events[idx]) final_output = torch.cat(results, dim=0)

这个场景的坑在于:汇总操作torch.cat/torch.stack可能在一个 chunk 还没算完时就开始执行,因为主流没有等待。而且chunk.record_stream(streams[idx])不能漏:chunk 的数据是在默认流上创建的,被送到 streams[idx] 上使用之后,要保护到该流算完。

另一个容易被忽略的细节:results[idx]是在 streams[idx] 上创建的,之后要在主流上使用,主流也要等待对应事件;而且results[idx]这个张量如果在下一次循环被覆写,同样要记得record_stream主流。我一般把创建流、使用流、回收保护流三者统一写清楚,宁可多写一行,也不赌调度器。

5.3 分叉/汇合与反向传播

反向传播场景里也经常出现隐形的多流问题。比如一个自研的多流并行模块,forward 时用了副流,backward 时虽然张量的 autograd graph 会保留,但中间张量如果生命周期很短,被分配器提前回收,backward 时就可能读到错误数据。

我的处理原则:凡是张量生命周期跨了流界限,就必须同时考虑 wait 和 record_stream。反向传播里有个好消息是 autograd 引擎会在计算图执行时做必要的依赖管理,但它是按张量的创建流来排的。如果你的自定义算子把一个新流上产生的张量交给另一个流的 backward kernel 用,还是需要显式同步。

实际中我建议避免在自定义autograd.Function里搞复杂的多流状态。如果确实要做,就在 forward 结束时用torch.cuda.current_stream().wait_stream(副流)把顺序钉死,并在中间张量上调用record_stream(当前流)。否则调试成本远大于那点性能收益。

5.4 避坑速查表

场景必须做的同步常见漏点
张量在流 A 创建,在流 B 消费B 等待 A 的事件,A 上记录事件忘了在 B 开头 wait
流 A 张量引用在 B 使用后释放tensor.record_stream(B)只 wait 不 record,内存被复用
CPU pinned → GPU 拷贝拷贝流的 non_blocking + 目标流 wait把 non_blocking 当成同步
多流并行 + 主流汇总主流等待所有分支流只等待部分流,或顺序错
循环内多流每轮重新建立依赖跨轮事件复用,依赖混乱

这张表基本覆盖了我这几年遇到的多流问题。新项目里一旦有人喊“多流后结果不对”,我都先让他对着表自查。

6. 排查技巧与工具

6.1 如何检查流是否按预期执行

排查流同步问题,我常用的探针是事件状态查询:

ev = torch.cuda.Event() with torch.cuda.stream(s1): ev.record(s1) print("event completed:", ev.query())

query()返回True表示事件在 host 视角已经完成。如果你发现目标流开始消费之前,事件还处于未完成状态,说明顺序依赖没建立。

另一个实用的办法是在目标流里放一个探针 kernel。不过实操里最有效的还是“最小化复现 + 二分对照”,事件查询只是辅助确认。我用事件报告耗时的时候也比较多,可以用ev1.report_elapsed_time(ev2)看两个事件中间的真实 GPU 时间差,帮助判断是不是某段执行被意外拖住了。

6.2 如何检查内存是否被复用

内存复用导致的隐形 bug 最难查,因为它不报错,只让数值偶尔不对。要确认是不是内存复用,可以:

  • 临时关闭缓存分配器,强制每次都真分配,比如PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,或用cudaMallocAsync后端。这只能辅助判断,正式跑不能这么干。
  • 在张量释放后立刻分配另一个张量,检查两个张量的data_ptr是否落在同一块地址。如果你发现新分配张量的 data_ptr 和刚释放的张量完全一样,并且旧流上的 kernel 还没跑完,基本就是复用了。
  • 用torch.cuda.memory_snapshot()查看分配器内部的 segment 和 block 信息,定位哪些块还挂着事件等待。这个工具对排查显存异常增长特别有效。

我实际项目里遇到过一种显存异常增长,就是靠memory_snapshot()定位到某些张量因为record_stream调用太密,事件没及时清理,导致块积压在延迟释放队列里。这属于另一类问题:record_stream多到“过度保护”,同样会拖累回收,要在该保护的那一小段和过度保护之间找平衡。

6.3 关于 wait_event 的几个细节

  • stream.wait_event(ev)和event.wait()都表示“在指定流上插入等待”,区别只在于调用的主语是谁。两个都要求ev已经被记录在某条流上,否则行为未定义。
  • torch.cuda.Stream.wait_stream(other)是快速写法,等价于让当前流等待 other 当前排队的所有工作。如果 other 后续还有新工作排入,不影响已经建立的等待。
  • 事件不要跨迭代复用太乱。我喜欢在每个迭代里新建事件,虽然有一点创建开销,但生命周期管理简单,避免顺序逻辑纠缠。

最后分享一个小习惯:写完多流代码后,我先跑一个带torch.cuda.synchronize()的快速验证,确认结果正确;然后再把synchronize()去掉,用并行版本验证性能和稳定性。两次都通过,才算真正安全。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询