GPU推理引擎抢占式调度:保障高优先级请求,优化端到端延迟
在连续批处理模式跑起来之后调度器看起来已经像模像样了。但一旦把服务真实放到多用户环境里我很快就被一个问题卡住低优先级的长请求霸占了GPU不肯走高优先级的交互式请求在队列里干等端到端延迟直接飙到几秒甚至十几秒。怎么让高优先级请求“插队”拿到资源又不把低优先级请求彻底饿死这其实就是推理引擎调度系统里非常硬核的一个话题——抢占Preemption。花了两个周末把这块逻辑重写了一遍之后我把它完整记录下来。这篇文章主要面向已经实现了基础连续批处理调度、想进一步优化服务质量QoS和优先级控制的同行也适合对大模型推理服务端架构感兴趣、想理解“为什么线上推理服务能保证响应延迟”的读者。1. 抢占不是“杀掉任务”而是“暂停-恢复”的资源再分配1.1 为什么连续批处理跑通了还不够先说背景。我的引擎按照经典的连续批处理思路实现了调度器核心每个step从等待队列和运行队列里挑选request组成一个iteration batchforward一次然后根据状态把request在各个队列之间迁移。最初版本只有两级状态——Running和Waiting调度逻辑非常直白def schedule(self): running self.running_queue waiting self.waiting_queue batch [] for seq in running: batch.append(seq) # 还有空间就从waiting补 free self.max_batch_size - len(batch) while free 0 and len(waiting) 0: seq waiting.popleft() batch.append(seq) free - 1 return batch这个版本在单用户、低并发场景下完全够用。问题出在混合负载。假设两个request同时到达请求A是长文档总结需要decode 2000个token请求B是单轮问答只需要decode 30个token。如果A先进队B就只能等A的所有decode token生成完毕才能被调度。这还不是最坏的情况。更糟糕的是当一个或者多个长请求占据了整个batch之后排在后面的短请求的排队延迟会变得完全不可控。连续批处理确实提升了吞吐但它没有解决“重要请求被不重要请求堵住”的根本矛盾。1.2 抢占的本质GPU是单一大管家要理解抢占先得理解为什么GPU不适合做细粒度的任务打断。CPU上有完整的中断机制、时间片轮转操作系统可以在任意指令边界保存和恢复进程上下文。GPU不是这么玩的——GPU的核心数量巨大kernel launch又有固定开销如果在一个iteration中途把某个request踢出去硬件层面根本做不到只“暂停”其中的一部分计算。所以在推理引擎里抢占的单位不是token也不是算子而是“iteration”。说得更直白一点调度器只会在一个完整的forward step结束之后检查所有request的状态如果发现需要抢占某个request它就中止这个request继续进入下一轮的decode把它从Running队列挪到Waiting队列同时释放一部分显存资源。这个“在step边界做决策”的思路是整套抢占机制的基石。这和虚拟机抢占USB设备倒是有点神似——资源就一份谁抢到了谁用操作系统只能在合适的时机做切换。GPU调度同样如此只是把“设备切换”换成了“KV Cache保留与恢复”。1.3 抢占与任务取消的本质区别抢占很容易和取消混淆。取消是彻底放弃一个请求释放它占有的全部显存客户端收到错误或部分结果。而抢占是临时中止保留这个request在引擎内部的状态等资源充足时再恢复。这意味着被抢占的request需要回答三个问题它的KV Cache还在不在如果在它的显存空间有没有被释放它的生成进度记录在哪里恢复时会不会重复生成已经生成过的token谁来唤醒它靠新请求触发还是靠轮询这三个问题决定了抢占机制的具体实现方式。如果KV Cache完全释放恢复时就得从头prefill对于长上下文来说代价极其高昂。所以工程上几乎都采用“保留KV Cache只释放增量空间”的做法严格来说是“暂停-恢复”不是“重新计算”。2. 抢占方案选型从头设计还是流式接口2.1 两种风格Task级抢占 vs Request级抢占先说说这里的一个关键概念区分。在NCCL分布式训练或者Python多进程任务池里我们经常说抢占一个task基本上是整块任务的调度。而在推理引擎里请求天然是流式的——一个request的生命周期内要经历很多次forward。所以抢占的最小粒度应该是request级别而非task级别。也就是说我们不是“杀掉一个进程”而是“把一个request从运行态切回等待态”。这个区分非常关键因为它直接影响恢复逻辑任务级抢占一般需要重新执行整个任务请求级抢占只需要从断点继续decode前提是断点的KV Cache还在。2.2 三种实现路径对比在动手之前我对比了三条实现路径方案核心思路优点缺点Tick轮询每个step后扫描所有request的优先级实现简单延迟可控唤醒不及时调度粒度粗事件驱动高优先级请求入队时立即触发抢占检查抢占实时性好调度器逻辑复杂度高容易出现并发问题混合常规step边界检查 关键路径事件触发兼顾实现难度与实时性需要处理事件队列的优先级反转问题我最终选择了混合方案。引擎结构本身是单线程事件循环每个iteration之间有一个完整的调度窗口。在这个窗口里做优先级检查和抢占决策不会引入任何并发问题。而且高优先级请求入队时只需要设置一个“调度器脏标记”下一个iteration就能立刻感知实时性损失在一个step以内。这里顺带一提我在实验过程中踩过的一个坑千万不要试图在Pytorch的CUDA stream回调里做调度决策。GPU kernel执行完毕时CPU端确实会收到回调但在这个时间点切出请求CUDA context的状态处理非常容易出问题而且回调线程和主线程之间的同步会让性能倒退一截。老老实实在所有iteration结束后的CPU调度窗口里做才是正道。2.3 这与模型架构强绑定实现抢占还需要考虑模型本身是否支持灵活batch。如果你用的是PageAttention这类带显存分页管理的后端抢占实现会非常优雅——只需要把被抢占的seq的KV Cache对应的物理块标记为“非活跃”新请求可以使用这些块恢复时再重新映射。如果你用的是连续显存分配比如把所有request的KV Cache拼在一个大tensor里抢占就麻烦得多因为一旦把被抢占seq的KV Cache从大tensor里抠出去后面的数据都要往前挪恢复时还得重新插回去。我在引擎里用的是PageAttention思想的分块KV Cache每个逻辑块固定大小物理块通过block table映射。这让抢占的实现简单了一个数量级被抢占时只需要把block table从“活跃表”挪到“暂停表”不需要任何显存拷贝。3. 抢占调度器的核心实现3.1 数据结构设计给每个请求加上状态机抢占机制要求每个request都要有明确的优先级和状态标记。我在引擎里给Sequence加了一个枚举状态把原来的Running/Waiting细分成了四个class SeqStatus(enum.Enum): WAITING 1 # 在等待队列还没prefill RUNNING 2 # 正在decode/prefill中 SUSPENDED 3 # 被抢占KV Cache保留 FINISHED 4 # 正常结束并且在request对象上增加了两个字段priority整数值值越小优先级越高。这个由上游服务传入比如交互式请求传0离线任务传5。last_progress最近一次decode完成的step数用于恢复时校验进度。不需要额外维护复杂的数据结构一个优先队列就能同时服务调度与抢占。我用的是heapq最小堆按priority排序同优先级按到达时间排序。dataclass(orderTrue) class ScheduledItem: priority: int enqueue_time: float seq: Any field(compareFalse) running_heap: List[ScheduledItem] [] waiting_heap: List[ScheduledItem] [] suspended_heap: List[ScheduledItem] []三个堆分别对应三种状态。抢占发生时只需要把item从running_heap里取出来修改status然后放进suspended_heap。恢复时反过来。关键在于这个操作时间复杂度是O(log n)不会成为调度瓶颈。3.2 抢占决策逻辑多了一个维度一个批次的调度决策不能只看优先级因为优先级再高如果显存不够、KV Cache无法分配强行调度只会导致OOM。所以完整的调度算法变成了这样每个iteration结束时 1. 把FINISHED的seq从running堆移除 2. 检查suspended堆里是否有可恢复的请求高优先级 3. 检查waiting堆里是否有高优先级新请求 4. 计算当前批次剩余容量和显存余量 5. 如果高优请求需要资源而当前容量不足触发抢占 6. 当低优请求被抢占后把高优请求调度进batch 7. 正常调度剩余waiting请求核心代码我简化后是这样的def schedule(self): self._cleanup_finished() # 需要调度的候选请求按优先级排序 candidates [] # 先看suspended被抢占过的请求优先恢复再看waiting for item in self.suspended_heap: candidates.append(item) for item in self.waiting_heap: candidates.append(item) heapq.heapify(candidates) batch [] for item in candidates: if len(batch) self.max_batch_size: break seq item.seq if seq.status SeqStatus.SUSPENDED: # 恢复时重新挂载KV Cache不需要re-prefill self.block_table.restore(seq.block_table) seq.status SeqStatus.RUNNING elif seq.status SeqStatus.WAITING: # 首次调度执行prefill seq.status SeqStatus.RUNNING batch.append(seq) # 如果高优请求仍然因为batch满进不来就抢占低优请求 if self._has_high_priority_item(candidates, batch): self._preempt_low_priority(batch) return batch这个版本的调度逻辑天然实现了“被抢占过的请求优先恢复”。原因也很现实既然已经抢占了它一次说明它曾经拥有过资源如果恢复时还把它排在后面很可能导致它在suspended堆里长时间出不来造成饥饿。3.3 抢占执行三步走_preempt_low_priority函数是最核心的一段逻辑。它做的事情非常收敛从当前batch末尾开始往前扫描找到优先级最低而且还有pending状态的seq把它摘下来挂起。为什么从末尾开始因为尾部通常是这个iteration里最后加入的请求它的prefill/decoded进度最少回滚代价最小。def _preempt_low_priority(self, batch): # 逆序扫描batch找到最低优的RUNNING请求 for seq in reversed(batch): if seq.priority self.preempt_threshold and seq.status SeqStatus.RUNNING: # 1. 从block_table中移除物理块的活跃标记 self.block_table.suspend(seq.block_table) # 2. 标记状态 seq.status SeqStatus.SUSPENDED seq.last_progress seq.output_ids_len # 3. 从running堆删除惰性删除 heapq.heappush(self.suspended_heap, ScheduledItem(priorityseq.priority, enqueue_timetime.time(), seqseq)) # 腾出的空间立即被高优请求占用结束抢占即可 break # 重新调度腾出的空间 self._fill_batch_from_waiting(batch)需要注意两个实现细节。第一从running堆里删除item时我用了惰性删除因为heapq没有内置remove操作直接在堆顶标记seq.status变更即可下次调度时如果从running堆里pop出已经SUSPENDED状态的item直接跳过。这比手动维护额外的dict做索引要简单得多。第二suspend操作只做逻辑标记不动显存里的实际数据。block_table里存的是物理块号到逻辑块号的映射suspend时做的事情是把映射关系拷贝一份到seq的专属表然后在全局block_table里把这些物理块标成可复用。恢复时再把映射关系重新注册回去。整个过程不涉及cudaMemcpy因此抢占本身的时间成本极低实测只有几十微秒。3.4 batch怎么重新布局抢占最容易被忽略的坑是batch tensor的重新布局。在连续批处理里input_ids、position_ids、attention_mask这些tensor的shape是[batch_size, seq_len]batch_size是动态的。当抢占发生后batch里少了一个seq多了一个seq如果直接按照新的request列表重新stack一遍被抢占seq在显存里的KV Cache块虽然标记为可复用但物理块在没有被立即覆盖前数据仍然在。新seq的prefill计算完全不依赖被抢占seq的数据所以只要prepare_inputs时按当前batch重新打包tensor即可。我最初犯过一个错误为了减少tensor拷贝把被抢占seq的KV Cache物理块直接清零结果不但多花了大量无用时间还拉低了显存带宽。正确的做法是懒清理——物理块在真正分配给新seq时才会被覆盖写入不需要手动清零。大模型推理引擎里一切以带宽为重能省的显存操作都要省。4. 抢占触发的边界条件与性能实测4.1 什么时候该抢占什么时候不该抢调度器不是看到高优请求就立刻抢占。在真实服务里如果每个低优请求都被频繁抢占它们可能永远无法完成这就是饥饿问题。我的实现里设置了一个preempt_threshold参数只有当低优请求的优先级高于这个阈值时才允许被抢占。比如交互式请求priority0后台任务priority5preempt_threshold3。这样后台任务不会抢交互式的资源但交互式可以抢后台任务的。还有两个附加条件我强烈建议加上被抢占的seq的已生成token数不能太多。如果它已经decode了几百个token再等一小会儿就完成了此时抢占反而浪费算力。我在代码里判断seq.output_ids_len self.max_preempt_len才允许抢占。被抢占的seq必须处于decode阶段。如果它还在prefill阶段意味着KV Cache都还没建好抢占它更划算因为恢复时也不需要re-prefillKV Cache没建好就在恢复时重新prefill即可。4.2 实测数据在OPT-1.3B上的表现我用OPT-1.3B模型、A10 GPU做了三组对比实验。batch size固定16模拟两个场景场景A是30个低优长请求各需要decode 512 token场景B是30个高优短请求各需要decode 32 token。三组测试分别是无抢占、抢占阈值3、抢占阈值0即任何低优请求都可能被抢占。配置高优请求平均延迟低优请求平均延迟低优请求完成率无抢占3.42s1.81s100%抢占阈值30.67s2.74s96.7%抢占阈值00.52s4.13s83.3%数据非常直观。无抢占时高优请求被长请求堵死延迟受到严重拖累。加了抢占阈值3之后高优请求延迟下降了约80%低优请求虽然变慢了一点但没有出现大面积饿死的情况。抢占阈值0时虽然高优延迟进一步下降但低优请求有将近17%被饿死显然不可接受。我个人把阈值3作为默认值。在真实线上环境还可以为上层的“租户”概念增加优先级分层在调度器内部做两级抢占——跨租户抢占和租户内抢占。这样既保证高优租户的SLO也不至于因为个别高优请求反复抢占同一低优请求。4.3 一个容易忽略的显存陷阱实现抢占时KV Cache的持有方式有一个微妙但致命的坑。假设引擎同时跑16个请求每个请求的KV Cache都分配了固定长度的物理块。某时刻抢占了其中4个显存释放了约25%。但这个释放出来的显存是碎片化的——每个被抢占请求都分布在不同的物理块上这些块之间不连续。如果后端管理器按“连续大块”方式管理显存这25%的空闲根本没法被利用。所以我在引擎里干脆全部采用分页式分配物理块大小固定以块为单位挂载和卸载。这样即使释放出来的块分散在显存各处新请求的KV Cache也能按块分配拼接起来用。这也是为什么前面说“显存分页管理是抢占的基础设施”。再补充一点被抢占的seq的KV Cache虽然逻辑上被释放了但物理上的数据在真正被覆盖之前仍然存在。如果引擎崩溃后需要重启恢复这个数据已经不可靠所以不要用它做任何持久化假设。它只是为了让抢占恢复时不丢失进度的临时缓冲。5. 抢占之后的恢复路径与状态一致性5.1 恢复流程不重新prefill一旦资源充足比如batch里某请求完成或者高优请求已经输出完毕被抢占的seq从suspended_heap里被pop出来恢复到running队列。恢复时的核心原则是“不重新prefill”。什么意思应当这样理解被抢占那一刻seq已经生成了部分token。假设它已经decode了100个token它的完整输入长度是“prompt 100个生成的token”。恢复时如果走普通的调度流程需要对这整段输入重新计算KV Cache成本极高。但如果KV Cache分页保留只需要把block_table重新映射到物理块然后从第101个token开始继续decode。这一步几乎不花时间在CPU侧只是几十次指针操作。代码里恢复逻辑非常简单def restore_seq(self, seq): # 重新把block_table挂回全局页表 self.page_manager.attach(seq.block_table) seq.status SeqStatus.RUNNING seq.last_progress seq.output_ids_len # 下一次iteration会从output_ids_len的位置继续decode真正的难点不在于恢复本身而在于如何保证状态一致性。被抢占的seq可能在SUSPENDED状态期间上游客户端已经等不及断开了连接。如果恢复后继续生成然后把结果发给一个已经关闭的socket轻则浪费算力重则引发回调空指针崩溃。因此在恢复调度时需要先检查seq对应的client连接是否仍然存活。这个检查在真实系统里往往被忽略但我认为它是必须的。5.2 被抢占请求的计时问题大多数推理引擎都会给请求设置超时时间。在无抢占系统里超时从请求到达服务器开始计时到完整响应返回结束。在引入抢占机制后计时会带来两类问题。第一类被抢占期间的时间算不算超时我的做法是算但用独立的“初试时间戳”来统一管理。也就是从request第一次被调度第一次prefill开始到最终完成的时间会比纯decode时间更长但这个时长也包含了被抢占的等待时间。超时时间设置为“预期decode耗时 x 1.5 最大允许排队时间”这样即使偶尔被抢占也不会误杀请求。第二类被抢占后重新进入等待队列怎样保持排队顺序这里我特意没有把enqueue_time重置为当前时间。如果重置被抢占的请求就会“排队到队尾”在持续高优负载下永远排不上号。正确做法是保留原始的enqueue_time抢占恢复时按这个时间戳排序这样它依然排在同等优先级请求的前面最大程度保证公平性。5.3 多层调度下的优先级反转如果你的引擎前面有网关层网关层也在做队列管理那么网关和引擎之间的优先级字段必须透传。否则可能出现这种情况网关层把高优请求排在低优请求前面但引擎内部调度时两个请求的priority字段都是默认值0完全分不出来。这其实是分布式系统里经典的优先级反转问题在推理服务里的变体。我在项目里做了两处修改来规避这个问题。第一网关传入的HTTP header里显式携带X-Priority引擎解析后转为内部整数值。第二引擎内部调度时以内部优先级为主不信任任何客户端自定义的优先级字段防止恶意请求把自己标记为最高优先级从而垄断GPU。这两点虽小但在线上环境中作用巨大。6. 可观测性抢占不可怕可怕的是盲目抢占6.1 抢占指标该统计什么加了抢占之后我立刻发现一个效率问题——没有可观测性很难回答“我的抢占策略到底有没有效”。所以我加了一组指标全部基于Prometheus的Counter和Histogram生成指标说明preempt_total抢占总次数preempt_by_priority按优先级维度统计的抢占次数preempt_recover_total恢复总次数preempt_wasted_tokens被抢占时已经生成但永远不会被使用的token数request_e2e_latency_seconds请求端到端延迟直方图request_preempt_wait_seconds请求在SUSPENDED状态下等待的时长这里面我觉得最有用的是preempt_wasted_tokens。它的计算方式很简单被抢占seq已经生成的token数量。如果这个值一直很大说明你的抢占决策做得太晚了——很多请求都快生成了才被抢占白白浪费算力。6.2 黑盒测试如何确认抢占没有破坏正确性抢占逻辑影响的是调度层的状态机不会改变模型本身的数值计算路径。只要恢复时KV Cache挂载正确模型输出应该和没有被抢占时完全一致。为了验证这一点我做了一个非常朴素的黑盒测试构造同一组prompt分别在“无抢占”和“强制抢占”两种模式下跑推理然后逐token对比生成结果。具体做法是让低优请求每生成50个token就强制被抢占一次恢复后再继续。理论上只要KV Cache挂载没有丢tokendecode时取了正确的position_ids和KV Cache生成结果是确定性的两次运行应完全一致。实测结果也确实如此两次生成的前几百个token完全一致没有观察到任何偏差。6.3 调试抢占逻辑的设施模拟器最后强烈推荐一个思路在真实GPU上调试调度逻辑太慢了一次iteration几毫秒连续跑几百个iteration才能复现一个问题。我写了一个纯CPU的调度模拟器直接喂给调度器假的request和假的KV Cache分配/释放回调。这样可以在毫秒级内跑完上千轮调度快速验证各种边界条件比如同时到达10个高优请求、低优请求被抢占时正好完成了95%等。模拟器本身不到两百行用unittest.mock替换掉模型forward和显存管理即可。我在模拟器上发现了两个真实代码里极难发现的bug——其中一个是在抢占恢复后忘了更新position_ids导致恢复后token重复生成。这种问题如果在GPU上调试可能要跑几个小时才能从海量日志里定位但在模拟器里几秒钟就暴露了。7. 一些实战经验总结本轮写抢占部分我个人最深的感触是推理引擎的调度本质上是在资源有限的前提下做取舍。抢占机制让调度器有能力在高优请求到来时快速腾出资源但它不是银弹。如果非要把所有墙都推倒重来你应该把以下三条原则焊死在代码注释里抢占必须发生在iteration边界不要在算子执行中打断。被抢占的请求要保留KV Cache的物理块映射恢复时不要重新prefill。抢占决策必须同时考虑优先级、生成进度和显存碎片不要只按优先级一刀切。在代码之外也有一个让我印象很深的教训抢占逻辑的单元测试不复杂难的是模拟真实流量下的调度模式。我建议每一个想把抢占做进生产环境的团队都至少维护一个基于trace回放的调度回归测试——把线上真实请求的时间戳和优先级录制下来用模拟器回放。这样可以确保你重构抢占逻辑时不会一夜之间把线上服务质量调回解放前。