ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

多流推理必读:record_stream 和 wait_event 的职责与正确用法

多流推理必读:record_stream 和 wait_event 的职责与正确用法 一次线上推理服务的延迟优化把原本跑在默认流上的处理流程拆成了三路 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_streamCPU 预处理 拷贝、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_blockingTrue) # 模型推理 with torch.cuda.stream(compute_stream): outputs model(inputs_gpu) # 结果回传 with torch.cuda.stream(copy_back_stream): results outputs.to(cpu, non_blockingTrue)当时我天真地以为进入不同 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, devicecuda) with torch.cuda.stream(s2): b a * 2.0 if not torch.isfinite(b).all(): print(fi{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, devicecuda) 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_blockingTrue) 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_blockingTrue) 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_blockingTrue) # 计算时切换回主流但主流必须等待 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, dim0)这个场景的坑在于汇总操作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_CONFexpandable_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()去掉用并行版本验证性能和稳定性。两次都通过才算真正安全。
返回列表