ARTICLE DETAIL

资讯详情

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

从零手搓AI工程:推理调度与显存管理实战

从零手搓AI工程:推理调度与显存管理实战 1. 从零手搓AI工程为什么我不建议你直接调包很多人一上来就想跑通一个能对话的模型或者直接拉一个开源仓库改改就上线。我见过太多团队模型效果在Demo阶段惊艳一进生产环境就崩——延迟飙到几秒、显存动不动爆掉、并发一上来就排队。问题出在哪不是模型不行是AI工程这层没搭好。ai-engineering-from-scratch这个标题核心不在“AI”而在“engineering”和“from scratch”。它讲的不是怎么调API而是从最底层把一套AI系统该有的东西自己搭出来数据怎么流、推理怎么调度、显存怎么管、服务怎么扩。适合谁看适合那些已经会写Python、跑过几个模型但一遇到“上线”就发怵的工程师也适合想真正理解推理框架内部在干什么、不想永远当调包侠的人。我自己走过这条路。最早做推理服务直接拿现成框架套结果一个batch size设错P99延迟从200ms跳到2s。后来逼着自己从零写了一遍调度逻辑才明白那些框架里的参数到底在权衡什么。这篇就把我从零搭AI工程链路时踩过的关键节点、做过的取舍、以及那些文档里不会写的经验完整拆一遍。2. 先想清楚从零搭AI工程到底在搭什么2.1 把AI工程拆成四层别一锅炖很多人说“搭AI系统”脑子里是一团浆糊。我的习惯是拆成四层每层职责分明数据层负责样本的读取、预处理、批处理。这层的关键是吞吐和顺序。训练时要做shuffle推理时往往要保序。模型层模型结构定义、权重加载、前向计算。这层的关键是显存布局和计算图。调度层请求怎么排队、怎么组batch、超时怎么处理。这层是AI工程和普通后端最大的区别所在。服务层对外暴露接口、健康检查、扩缩容。这层反而最接近传统后端。为什么这么拆因为每一层的瓶颈和优化手段完全不同。数据层卡I/O模型层卡显存调度层卡策略服务层卡网络。你如果不拆开出了问题根本不知道从哪下手。我见过一个服务QPS上不去团队一直在加机器最后发现是调度层用了全局锁加机器根本没用。2.2 “from scratch”不是让你重写CUDA这里要澄清一个误区。from scratch不是让你从汇编开始写也不是让你重写CUDA kernel。它的意思是不依赖黑盒框架自己把控制流和数据流串起来。你可以用PyTorch做张量计算用NumPy做数据处理但调度、组batch、显存复用这些逻辑得自己写。为什么因为只有自己写一遍你才知道框架帮你做了什么。比如动态batch框架里就是一个参数但自己实现时你要考虑新请求来了是等还是立刻发等的超时怎么设不同长度的序列怎么padding才不浪费这些决策直接影响成本和延迟。我自己的经验是自己写过一遍调度之后再用任何框架看参数都能猜到它内部大概怎么实现的调参不再是玄学。2.3 一个最小可用的AI工程骨架长什么样先给一个我常用的最小骨架后面所有讨论都围绕它展开# 伪代码展示核心结构 class InferenceEngine: def __init__(self, model, max_batch_size, max_seq_len): self.model model self.queue RequestQueue() self.scheduler BatchScheduler(max_batch_size, max_seq_len) self.memory_pool MemoryPool() def run(self): while True: batch self.scheduler.form_batch(self.queue) if batch: inputs self.preprocess(batch) outputs self.model(inputs) self.postprocess(outputs, batch)这个骨架里RequestQueue管请求排队BatchScheduler管组batchMemoryPool管显存复用。看起来简单但每个模块都有坑。下面逐个拆。3. 请求队列与动态组batch延迟和吞吐的拉锯战3.1 为什么不能来一个请求就推理一次最朴素的实现是一个请求进来直接调模型返回结果。这在单用户场景没问题但一旦并发上来GPU利用率会低得可怜。因为模型推理是计算密集型的单条请求的计算量往往喂不饱GPU。GPU大部分时间在等数据搬运而不是在算。动态组batch就是解决这个的。核心思想攒一小段时间的请求凑成一个batch一起推理。这样GPU一次算多条吞吐能提升几倍到几十倍。但代价是延迟——每个请求都要等一会儿才能被处理。这里有个关键权衡等多久。等太久延迟高等太短batch太小吞吐上不去。我的经验值是在线服务一般等5-20ms。这个数字怎么来的假设你的模型单条推理要50ms那么等10ms组batchbatch size到8总延迟是105060ms比单条50ms只多了10ms但吞吐翻了8倍。这笔账很划算。3.2 组batch时最容易忽略的padding浪费组batch有个隐蔽的坑序列长度不一致。比如一个batch里有的请求输入长度是10有的是500。如果直接padding到500那短序列的计算全是浪费。我实测过一个场景平均长度80最大长度512如果无脑padding到512有效计算只有15%左右GPU大部分算力在算padding。解决办法有两种。一种是分桶把长度相近的请求放一个batch。比如0-64一桶64-128一桶。这样padding浪费小。另一种是动态padding每个batch只padding到当前batch的最大长度而不是全局最大长度。这两种可以结合用。分桶的实现要注意桶的边界怎么定。定太细每个桶请求少组不成大batch定太粗padding浪费又上来了。我的做法是先统计线上请求的长度分布按分位数来定桶边界。比如P50、P80、P95各切一刀形成4个桶。这样大部分请求落在前几个桶里padding浪费可控。3.3 超时和优先级别让一个慢请求拖垮整个队列队列里最怕什么最怕一个超长请求卡住。比如一个请求输入长度是10000其他都是100。如果它和别的请求组一个batch整个batch都要padding到10000其他请求全被拖慢。我的处理方式是设长度上限。超过上限的请求直接拒绝或者走单独的低优先级通道。这个上限怎么定看你的业务。如果是对话场景一般512或1024就够了。如果是文档处理可能要到4096。但不管多少一定要有上限否则显存会被撑爆。另外队列要有优先级。比如健康检查请求、内部调用请求应该比外部用户请求优先级高。实现上可以用多个队列调度时先从高优先级队列取。但要注意高优先级队列不能饿死低优先级队列否则低优先级请求永远排不上。我一般给每个队列设一个配额比如每轮调度至少从低优先级取一个。4. 显存管理AI工程里最硬的骨头4.1 显存都去哪了算一笔账很多人对显存没概念觉得模型加载进去就行了。实际上显存占用分好几块占用项说明典型占比模型权重参数本身30-50%激活值前向计算的中间结果20-40%输入输出batch数据10-20%碎片分配释放产生的空洞5-15%模型权重是固定的但激活值和输入输出随batch size和序列长度变化。这就是为什么batch size不能无限大——激活值会线性增长。我见过一个模型权重只占4GB但batch size开到32时激活值占了12GB直接OOM。算这笔账的意义在于你要知道你的显存预算怎么分配。我的习惯是留20%的余量剩下的按权重、激活、输入输出分。如果激活值占比太高就要考虑用梯度检查点虽然推理用不上但训练时有用或者减小batch size。4.2 显存池为什么malloc会拖慢推理如果你每次推理都新分配显存推理完再释放会有两个问题。一是分配本身有开销虽然单次不大但高频调用会累积。二是碎片化反复分配释放不同大小的块显存会变得千疮百孔最后明明总量够却分配不出连续的大块。解决办法是显存池。启动时一次性申请一大块显存之后所有分配都从池子里切释放时还给池子不还给系统。这样分配释放都是O(1)而且没有碎片。实现显存池的关键是块大小。如果块太大小请求浪费如果块太小大请求要拼多个块管理复杂。我的做法是按2的幂次分块比如256MB、512MB、1GB、2GB。请求来了向上取整到最近的块大小。这样内部碎片最多50%但管理简单实测下来比精细管理更稳。4.3 显存复用的边界什么时候不能复用显存复用听起来很美但不是所有场景都能用。有个关键前提复用块的生命周期不能重叠。比如两个请求同时在线它们的输入输出不能共用一块显存。我踩过一个坑为了省显存把输入和输出的buffer复用了。结果模型是in-place操作输出直接覆盖了输入导致后续处理拿到的是脏数据。排查了半天才发现是复用边界没划清。所以复用要满足两个条件一是时间上不重叠前一个请求彻底处理完才能把块还给池子二是逻辑上不依赖输出不依赖输入的原值。第二条尤其要注意很多模型有残差连接输出依赖输入这种就不能复用。5. 推理调度让GPU一直忙起来5.1 同步推理 vs 异步推理选哪个同步推理就是发一个batch等结果再发下一个。异步推理是发一个batch不等结果继续发下一个结果通过回调或future返回。同步的好处是简单逻辑清晰。坏处是GPU会有空档——等结果的时候GPU在闲着。异步能填满空档但复杂度高要处理结果乱序、错误传播等问题。我的建议是如果单batch推理时间远大于调度开销用同步就够了。比如单batch要100ms调度只要1ms那同步的1%空档可以接受。但如果单batch只要5ms调度要1ms那20%的空档就值得用异步了。异步的实现可以用CUDA stream。每个batch绑一个stream多个stream可以并发。但要注意stream之间如果有依赖要加event同步。我一般用两个stream交替一个在算的时候另一个在准备数据这样能重叠计算和数据搬运。5.2 连续批处理一个被低估的优化连续批处理continuous batching是这两年推理优化的热点。传统组batch是“静态”的凑齐一个batch一起算算完一起返回。连续批处理是“动态”的batch里的请求算完一个就移出一个同时移入新请求。这样batch size始终是满的GPU利用率更高。实现连续批处理的关键是attention的mask。因为batch里不同请求的序列长度不同而且有的请求已经算完了有的还在算所以attention要能处理这种“部分完成”的状态。这需要改attention的实现不能直接用现成的。我实测下来连续批处理在长序列场景下提升明显能到2-3倍吞吐。但短序列场景提升有限因为请求很快就算完了移入移出的开销占比高。所以要不要上看你的业务场景。5.3 调度策略FCFS、优先级、还是公平调度调度策略决定了请求的处理顺序。常见的有三种FCFS先来先服务简单公平但一个慢请求会拖住后面所有请求。优先级重要请求先处理但低优先级可能饿死。公平调度每个用户或每个来源分一个配额保证都能得到服务。我的经验是混合用。默认FCFS但对超长请求降级对高优先级请求插队。具体实现可以用一个优先队列优先级由请求类型和等待时间共同决定。等待时间越长优先级越高这样能防止饿死。这里有个细节优先级不能只看类型还要看等待时间。否则高优先级请求源源不断低优先级永远排不上。我一般设一个老化因子每等100ms优先级提升一级。这样最坏情况下低优先级请求等几秒也能被处理。6. 服务化与压测从能跑到能扛6.1 接口设计别把内部结构暴露出去服务化第一步是定接口。我的原则是接口要稳定内部随便改。比如对外只暴露一个/predict接收文本返回结果。内部的batch size、调度策略、显存管理都不应该出现在接口里。但有一个例外超时时间。这个应该让调用方指定因为不同业务对延迟的容忍度不同。我一般设一个默认值比如5秒调用方可以覆盖。但覆盖有上限比如最多30秒防止有人设个无限大把服务拖死。另外接口要支持流式返回。对于生成式模型用户不想等全部生成完才看到结果。流式返回可以边生成边推体验好很多。实现上可以用SSE或WebSocket。SSE更简单单向推送就够了。6.2 压测怎么测才准压测最容易犯的错是用固定输入压。比如所有请求都是同样长度、同样内容。这样测出来的吞吐是虚高的因为padding浪费最小、缓存命中率最高。真实场景下输入长度是变化的内容也是变化的。我的做法是用真实流量回放。把线上请求录下来压测时按真实分布发。如果没有真实流量就构造一个长度分布比如按P50、P80、P95、P99各占一定比例。这样测出来的数字才有参考价值。压测还要看P99延迟不能只看平均。平均延迟可能很好看但P99可能爆表。我见过一个服务平均延迟50msP99到了3秒。原因是少数超长请求拖慢了整体。这种问题只有看P99才能发现。6.3 扩缩容什么时候加机器什么时候优化代码QPS上不去第一反应是加机器。但加机器之前先看GPU利用率。如果GPU利用率已经90%以上那加机器有用。如果GPU利用率只有30%那加机器是浪费问题在调度或数据层。我一般看三个指标GPU利用率、队列等待时间、batch size。如果GPU利用率低、队列等待时间长、batch size小说明调度有问题请求没凑成足够大的batch。这时候应该优化调度策略而不是加机器。如果GPU利用率高、队列等待时间长那说明确实算力不够该加机器了。但加机器之前还可以考虑模型量化或蒸馏把单次推理的计算量降下来。这比加机器更省钱。7. 那些文档里不会写的实操心得7.1 日志要记什么不记什么AI服务的日志和普通后端不一样。普通后端记请求响应就够了AI服务还要记batch信息。比如每个batch的size、序列长度分布、推理耗时。这些信息是调优的依据。但日志不能记太多否则I/O会成为瓶颈。我的做法是采样记。比如每100个batch记一条详细日志其余只记摘要。摘要包括batch size、平均长度、推理耗时。详细日志包括每个请求的ID、长度、等待时间。还有一个坑不要在推理线程里写日志。写日志是I/O操作会阻塞推理。我一般用一个单独的日志线程推理线程把日志丢进队列日志线程异步写。7.2 模型加载冷启动怎么优化服务重启时模型加载要时间。大模型加载可能要几十秒甚至几分钟。这期间服务不可用用户体验很差。优化手段有几个。一是预热服务启动后先用几个假请求跑一遍把显存池、CUDA context都初始化好。这样第一个真实请求不会特别慢。二是懒加载不是所有模型都一开始就加载按需加载。但懒加载有个问题第一次请求会特别慢。所以适合低频模型。三是模型分片把大模型切成几块加载一块就能提供部分服务。但这需要模型支持分片推理实现复杂。我的经验是预热最划算。实现简单效果明显。预热请求用真实分布的数据跑几十个就够了。7.3 错误处理模型推理失败了怎么办模型推理可能失败原因很多输入超长、显存不足、CUDA错误。失败之后怎么办直接返回500那用户体验很差。我的做法是分级处理。输入超长直接返回400告诉用户输入太长。显存不足先尝试减小batch size重试如果还不行返回503告诉用户服务繁忙。CUDA错误这个比较严重可能是硬件问题记录日志并返回500同时触发告警。重试要注意幂等性。推理一般是幂等的同样的输入应该得到同样的输出。但生成式模型有随机性重试可能得到不同结果。如果业务要求确定性要设随机种子。7.4 一个容易被忽略的细节时钟同步分布式部署时多个实例的时钟可能不同步。这会导致日志时间戳错乱排查问题时很痛苦。我踩过一次坑两个实例的日志时间差了3秒导致我以为请求先到了A再到B实际上是反的。解决办法是用NTP同步时钟并且在日志里记录相对时间比如从服务启动开始的毫秒数。这样即使绝对时间有偏差相对顺序是对的。8. 从零搭完之后我学到了什么自己从零搭一遍AI工程链路最大的收获不是代码本身而是对权衡的理解。以前调框架参数是试出来的现在调参数是算出来的。比如batch size设多少我会先算显存预算再看延迟要求最后定一个值。而不是试32不行试16试16不行试8。另一个收获是对瓶颈的敏感度。现在看到QPS上不去我会先看GPU利用率再看队列等待再看batch size基本能定位到是哪一层的问题。而不是盲目加机器或改代码。还有一个体会是简单方案往往更稳。我一开始想搞很复杂的调度策略结果bug一堆。后来退回到FCFS加超时反而稳定运行了很久。复杂策略不是不好而是要先证明简单方案不够用再上复杂的。最后分享一个小技巧给每个请求打一个trace ID从进队列到出结果全链路记录。这样排查问题时能完整看到请求在每一层的耗时。我靠这个定位过好几次性能问题比看聚合指标有用得多。
返回列表