ARTICLE DETAIL

资讯详情

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

从零手搓AI推理引擎:显存优化、KV Cache与并发调度实战

从零手搓AI推理引擎:显存优化、KV Cache与并发调度实战 1. 从零手搓AI工程为什么我不建议你直接调包很多人一上来就想跑通一个能对话的模型或者直接拿个开源框架套壳上线结果遇到显存溢出、推理延迟飙到几秒、并发一上来就崩。这些问题的根子不在模型本身而在于对AI工程链路的理解是断层的。ai-engineering-from-scratch这个方向核心不是让你重复造轮子而是让你亲手把轮子拆一遍知道每个齿轮咬合在哪里。它适合那些已经会用Python、了解基本深度学习概念但一遇到部署、优化、显存管理就发怵的开发者。我见过太多人能把模型训练脚本跑通却说不清楚KV Cache到底缓存了什么也搞不明白为什么batch size调大一点就OOM。这篇内容就是把这些黑盒一个个撬开从张量在内存里怎么摆放到推理引擎怎么调度请求全部用可复现的代码和实测数据讲清楚。你不需要买昂贵的多卡机器一张消费级显卡甚至CPU就能跟着走完大部分流程。读完你至少能明白为什么有些模型加载要几十秒有些却能秒级响应为什么同样的参数量推理速度能差出三倍。这些认知比你会调多少个API都值钱。1.1 先搞清楚“从零”到底指什么“从零”不是让你从晶体管开始造芯片也不是让你手写CUDA内核。在AI工程语境下它指的是不依赖高度封装的推理框架比如某些一行代码就能加载模型的库而是用PyTorch或NumPy这类基础工具把模型加载、权重映射、前向计算、内存管理这几个环节自己实现一遍。我试过用纯NumPy实现一个两层LSTM的推理虽然慢得离谱但那次之后我彻底理解了hidden state的传递逻辑。对于Transformer类模型从零意味着你要自己处理tokenizer的输出、自己构建attention mask、自己管理KV Cache的追加和截断。这些操作在封装库里都是自动的但一旦线上出现输出乱码或者显存泄漏你连排查的入口都找不到。所以这个“从零”的边界是底层算子可以用现成的但数据流和内存流必须自己掌控。1.2 一张消费级显卡能跑多远我拿RTX 3060 12GB做了一组实测。加载一个7B参数的模型如果用FP16精度光权重就占14GB直接爆显存。但换成INT8量化权重降到7GB左右加上KV Cache和中间激活勉强能跑起来生成速度大约每秒8到12个token。如果再降到INT4权重只有3.5GB速度能提到每秒20个token以上但输出质量会明显下降尤其是代码生成任务会出现重复和逻辑断裂。所以“从零”的硬件门槛并不高关键是你要知道每一步操作消耗了多少显存。我习惯在代码里插一个显存监控函数每执行一个阶段就打印一次torch.cuda.memory_allocated()这样能精确知道是权重加载占多了还是KV Cache没释放。这个习惯帮我省下了至少三次重买显卡的钱。1.3 为什么先跑通再优化是错的很多教程告诉你先把模型跑起来再慢慢调优。但在AI工程里这个顺序会让你后期付出巨大代价。比如你一开始用FP32加载模型跑通了然后想优化到INT8你会发现精度损失导致输出完全不可用这时候你回头改代码发现所有中间激活的数值范围都变了之前调好的temperature和top_p全部失效。正确的做法是一开始就用目标精度比如INT8来搭建推理链路哪怕速度慢一点也要保证数值一致性。我踩过这个坑一个文本分类任务FP32下准确率92%换成INT8后掉到78%排查了两天才发现是LayerNorm的epsilon在低精度下需要调整。所以从零构建时精度策略要前置不要留到最后。2. 张量内存布局那些让你OOM的隐形杀手显存溢出是AI工程里最常见的报错但很多人只知道减小batch size却不知道张量的内存布局才是关键。同样一个形状为(batch, seq_len, hidden_dim)的张量在内存里是连续存储还是跨步存储占用的显存可能差出30%。更隐蔽的是PyTorch的view和reshape操作在某些情况下会触发隐式拷贝你以为只是改了个形状实际上显存里多了一份完整数据。我在调试一个长文本推理任务时发现序列长度从512加到1024显存占用不是线性增长而是翻了四倍。后来用torch.cuda.memory_summary()才看到attention矩阵的中间结果因为内存不连续触发了多次临时分配。这一章就把这些坑一个个挖出来告诉你什么时候该用contiguous什么时候该用expand以及怎么用内存池来复用缓冲区。2.1 连续内存与跨步访问的代价在PyTorch里一个张量的内存布局由stride决定。比如你有一个形状为(2, 3)的张量如果它是连续的stride就是(3, 1)意味着跳3个元素到下一行。但如果你对它做转置形状变成(3, 2)stride变成(1, 3)这时候它就不再是连续内存了。很多算子比如矩阵乘法要求输入是连续的否则会先调用.contiguous()做一次拷贝。这个拷贝在显存里就是实打实的新分配。我实测过一个(1024, 4096)的FP16张量连续时占8MB转置后如果不拷贝算子内部会额外分配8MB的临时缓冲区。在多层Transformer里这种临时分配累积起来能吃掉好几个GB。所以我的经验是在构建attention模块时尽量让Q、K、V的布局保持一致避免在计算过程中频繁转置。如果必须转置提前用.contiguous()显式拷贝一次比让算子隐式拷贝更可控。2.2 KV Cache的显存账本KV Cache是自回归生成的核心优化但它也是显存大户。假设模型有32层每层有32个注意力头每个头的维度是128序列长度是2048batch size是1精度是FP16。那么KV Cache的总大小是2 * 32层 * 32头 * 2048长度 * 128维度 * 2字节 1.07GB。这还只是batch size为1的情况如果并发到8直接飙到8.5GB。很多人在做多轮对话时发现聊到第十轮就OOM就是因为KV Cache没有做截断或滑动窗口。我的做法是在生成每个token后检查当前序列长度是否超过预设阈值比如4096如果超过就把最前面的K和V丢掉只保留最近的窗口。这个操作在HuggingFace的past_key_values里需要手动处理不能指望库自动帮你做。另外KV Cache的存储顺序也很讲究按层存储还是按头存储会影响后续拼接的效率。我习惯按层存储因为每层计算是独立的按层取用更符合计算图的顺序。2.3 用内存池复用缓冲区如果你在推理循环里频繁创建和销毁张量PyTorch的缓存分配器虽然会复用一部分内存但碎片化仍然会导致OOM。我试过在一个循环里每步都torch.zeros()一个中间张量跑1000步后显存占用比初始高了2GB。后来改成预分配一个固定大小的缓冲区每次用copy_或切片赋值显存占用就稳定了。具体做法是在初始化阶段根据最大序列长度和batch size预先分配好attention分数矩阵、中间激活、输出logits的缓冲区。然后在推理时只往这些缓冲区里写数据不新建张量。这个技巧在部署小模型时特别有用能把显存峰值降低40%左右。注意预分配的缓冲区要放在正确的设备上并且要确保每次写入前清空旧数据否则会出现数值污染。3. 推理引擎的调度逻辑从单请求到并发单条请求跑通只是第一步真正的挑战在于并发。当10个用户同时发来请求每个请求的序列长度不同生成长度也不同你怎么调度如果简单地串行处理延迟会线性叠加如果无脑并行显存直接爆炸。这一章讲的是推理引擎的核心调度策略包括连续批处理、分页注意力、以及请求优先级管理。这些概念在vLLM、TensorRT-LLM里都有实现但如果你不理解背后的逻辑调参就是瞎猜。我会用简化的代码模拟一个调度器让你看到每个请求在时间轴上的分布以及显存是如何被动态分配的。3.1 连续批处理到底连续在哪里传统的静态批处理要求所有请求的输入长度和输出长度一致否则就要padding到最大长度浪费大量计算。连续批处理Continuous Batching的核心思想是每个请求独立管理自己的KV Cache调度器在每个生成步检查哪些请求已经完成把完成的请求踢出去把新来的请求加进来。这样GPU的利用率能保持在较高水平。我实测过一个场景8个请求输入长度从128到1024不等输出长度从64到512不等。静态批处理需要padding到1024输入和512输出总计算量是8 * 1024 * 512连续批处理的实际计算量只有sum(输入长度 * 输出长度)大约减少了60%。实现连续批处理的关键是维护一个请求队列和一个运行队列每个生成步从队列头部取新请求从运行队列尾部移除已完成请求。注意KV Cache的索引要跟着请求ID走不能混淆。3.2 分页注意力如何解决碎片化即使有了连续批处理KV Cache的显存分配仍然会产生碎片。因为每个请求的序列长度是动态增长的如果一开始就按最大长度分配浪费严重如果按需分配又会产生大量不连续的小块。分页注意力PagedAttention借鉴了操作系统的虚拟内存分页思想把KV Cache切成固定大小的块比如每块存16个token的K和V然后用一个块表来记录每个请求占用了哪些块。这样显存分配就是按块进行的块与块之间不需要连续。我模拟过一个场景16个并发请求平均序列长度512用分页注意力后显存利用率从45%提升到85%以上。实现分页注意力的难点在于块表的维护和注意力计算时的块索引。在计算attention时需要根据块表把分散的K和V块拼成逻辑上的连续序列。这个操作在CUDA层面有优化但在纯Python里模拟会比较慢适合用来理解原理。3.3 请求优先级与超时处理线上服务里不是所有请求都同等重要。比如一个实时对话请求和一个离线批量摘要请求前者对延迟敏感后者对吞吐敏感。如果混在一起调度实时请求会被离线请求拖慢。我的做法是给每个请求打上优先级标签调度器在每个生成步优先处理高优先级请求的KV Cache更新。同时设置超时阈值如果一个请求在队列里等待超过2秒还没开始生成就直接返回超时错误避免用户干等。这个策略在压测时效果明显实时请求的P99延迟从3.2秒降到了1.1秒代价是离线请求的吞吐下降了15%。这个取舍要根据业务场景来定没有绝对的最优解。4. 量化与精度省显存不降智的实操边界量化是AI工程里最诱人的技术因为它能直接把显存占用砍半甚至砍到四分之一。但量化也是一把双刃剑用不好会让模型输出变成乱码。这一章不讲量化理论只讲实操中怎么选精度、怎么校准、怎么验证。我会用同一个模型在FP16、INT8、INT4三种精度下跑一组标准测试把输出质量、速度、显存占用列成表格让你看到每个精度档位的真实表现。同时分享一个我常用的校准集构建方法不需要标注数据只用模型自己的输出就能完成校准。4.1 FP16、INT8、INT4的实测对比我拿一个7B参数的对话模型做了三组测试硬件是RTX 3060 12GB输入固定为“请用三句话解释什么是机器学习”生成长度限制为128个token。结果如下精度显存占用生成速度token/秒输出质量主观评分1-5首次加载时间FP1614.2GBOOM无法运行--INT87.8GB11.34.528秒INT44.1GB22.73.219秒INT8下输出基本流畅偶尔有轻微重复INT4下会出现“机器学习是机器学习是机器学习”这种循环需要调高repetition penalty才能缓解。所以我的建议是如果显存允许优先INT8如果必须INT4一定要配合后处理去重。另外首次加载时间差异主要来自权重量化转换INT4的转换反而比INT8快因为位宽更小但精度损失更大。4.2 校准集不需要标注数据量化的关键是找到合适的缩放因子scale和零点zero point这需要校准数据。很多人以为校准集必须是有标注的领域数据其实不然。我通常用模型自己生成的100到200条文本作为校准集覆盖不同的输入长度和主题。具体做法是准备20个不同的prompt每个prompt让FP16模型生成10条输出把这些输出拼起来作为校准集。然后用这个校准集去统计每一层激活值的分布计算最小值和最大值从而确定量化参数。这个方法在我测试的多个模型上都有效INT8下的精度损失控制在1%以内。注意校准集的长度要覆盖你实际推理时的序列长度范围如果校准集全是短文本长文本推理时会出现截断误差。4.3 逐层量化与混合精度不是所有层都适合量化。我发现在Transformer里attention的QKV投影层对量化比较敏感而FFN层相对鲁棒。所以可以采用混合精度QKV保持INT8FFN用INT4。这样整体显存占用比全INT8低20%左右而输出质量几乎不变。实现混合精度需要在加载权重时逐层判断给不同的层打上不同的量化配置。PyTorch的torch.quantization支持这种细粒度控制但需要手动指定每一层的qconfig。我写了一个简单的规则如果层的参数量超过100万就用INT8否则用INT4。这个规则在7B模型上效果不错显存降到了6.2GB速度提升到15.1 token/秒。你可以根据自己的模型结构调整这个阈值。5. 从零搭建一个最小推理服务前面讲的都是零件这一章把它们组装起来。我会用一个不到500行的Python脚本实现一个支持并发请求、动态批处理、KV Cache管理、INT8量化的最小推理服务。不依赖FastAPI或Flask只用Python标准库的socket和threading让你看清HTTP请求是怎么变成模型输入的。这个服务能同时处理4个并发请求每个请求独立维护对话历史显存占用稳定在8GB以内。代码会分成几个模块请求解析、调度器、模型执行器、响应生成。每个模块我都会解释为什么这样设计以及哪些地方可以替换成更高效的实现。5.1 请求解析与tokenizer的坑请求进来首先是JSON格式的{prompt: ..., max_tokens: 128}。解析用json.loads就行但tokenizer这里有个大坑不同模型的tokenizer对特殊字符的处理不一样。比如有些模型会把换行符当成独立token有些会合并。如果你在拼接多轮对话时直接字符串相加tokenizer可能会在边界处产生意外的token。我的做法是用tokenizer的apply_chat_template方法如果模型支持或者手动在每轮对话后加上分隔符并确保分隔符在tokenizer的词表里是单个token。另外tokenizer的输出要转成input_ids和attention_mask后者在批处理时特别重要因为padding的位置必须被mask掉否则attention会关注到无意义的填充。5.2 调度器的线程安全实现调度器运行在一个独立线程里维护两个队列等待队列和运行队列。等待队列用queue.Queue线程安全运行队列用列表但访问时需要加锁。每个生成步调度器从等待队列取新请求非阻塞检查运行队列里哪些请求已经生成了max_tokens或遇到了结束符把它们移除并发送响应。然后对运行队列里的所有请求执行一次模型前向更新各自的KV Cache。这里的关键是模型前向必须在一个批次里完成所以要把所有运行请求的输入拼成一个batch。拼接时要注意序列长度对齐短的请求在前面补padding并用attention_mask标记。我实测过4个并发请求下调度器的开销不到5毫秒主要时间花在模型前向上。5.3 响应生成与流式输出生成响应时如果等所有token都生成完再返回用户会感觉卡顿。所以最好支持流式输出每生成一个token就通过socket发回去。但流式输出对调度器有额外要求每个请求的响应要独立维护一个缓冲区不能混在一起。我的做法是给每个请求分配一个response_buffer列表每生成一个token就追加进去然后立即通过socket发送。客户端收到的是一个个JSON对象每个对象包含当前生成的token和是否结束的标志。注意流式输出时如果客户端断开连接调度器要能检测到并清理对应的请求否则KV Cache会一直占着显存。我通过设置socket的超时和捕获BrokenPipeError来实现这个清理逻辑。5.4 压测与调优的实测数据我用wrk工具对这个最小服务做了压测4个并发连接每个连接发送10个请求输入长度128输出长度64。结果如下平均延迟1.8秒P99延迟3.4秒吞吐量2.2请求/秒。显存占用稳定在7.9GB没有泄漏。瓶颈主要在模型前向占了总时间的85%以上。如果把INT8换成INT4吞吐量能提到3.5请求/秒但输出质量下降明显。所以对于这个硬件配置INT8是甜点。另外我发现把max_tokens从64降到32延迟能减少40%因为KV Cache的追加操作少了。如果你的业务允许短回复这是一个立竿见影的优化。6. 踩过的坑与排查链路这一章不教新东西只复盘我踩过的三个典型坑。每个坑我都会给出完整的排查链路从现象到假设从验证到修复。这些坑在官方文档里找不到因为它们是工程实践中的边缘情况。比如显存明明够却报OOM比如模型输出突然变成乱码比如并发一高就卡死。这些问题的根因往往不在模型本身而在数据流或资源管理的某个角落。6.1 显存够却OOM碎片化的隐蔽性有一次我加载了一个INT8模型nvidia-smi显示显存占用7.2GB总显存12GB按理说还有4.8GB余量。但一跑推理就OOM。我用torch.cuda.memory_summary()打印了详细分配情况发现虽然总占用7.2GB但最大的连续空闲块只有1.1GB。原因是之前加载FP16模型时留下的缓存碎片没有释放。PyTorch的缓存分配器会保留已分配的内存块即使张量被释放了内存也不会还给系统。解决办法是调用torch.cuda.empty_cache()但这只能释放未使用的缓存块不能合并碎片。更彻底的做法是在加载新模型前重启进程或者用PYTORCH_CUDA_ALLOC_CONF环境变量设置max_split_size_mb限制内存块的最大分割尺寸减少碎片。我设置成128MB后同样场景下最大连续空闲块提升到了3.2GBOOM消失了。6.2 输出乱码tokenizer与padding的冲突有一次线上服务突然返回乱码输入是正常的中文输出却是“的的”。排查发现问题出在padding token上。我用的tokenizer的padding token ID是0而0在词表里对应的是unk。当batch里有短请求时我在前面补了0但attention_mask没有正确设置导致模型把padding位置也当成有效输入生成了大量unk。修复方法是确保attention_mask在padding位置为0并且在计算loss或生成时把padding位置的logits屏蔽掉。另外如果tokenizer没有定义padding token要手动设置一个不冲突的ID比如词表末尾新增一个pad。这个坑让我损失了一个下午但从此以后我每次加载tokenizer都会检查pad_token_id和eos_token_id是否合理。6.3 并发卡死GIL与CUDA流的交互最后一个坑最隐蔽。我用多线程实现并发请求每个线程独立调用模型前向。单请求时正常两个并发就卡死。用py-spy抓取堆栈发现两个线程都卡在torch.nn.functional.linear里。原因是PyTorch的CUDA操作默认在默认流上执行多线程同时提交CUDA操作时如果没有显式同步会导致死锁。解决办法是给每个线程分配独立的CUDA流或者更简单用一个全局锁把模型前向串行化。虽然串行化会降低并发性能但至少不会卡死。我选择了后者因为我的场景并发不高串行化的延迟增加可以接受。如果你需要真正的并行可以用torch.cuda.Stream()为每个线程创建独立的流并在流之间用事件同步。这个坑的教训是Python的GIL和CUDA的异步执行叠加在一起行为非常反直觉多线程推理要慎用。6.4 排查工具链的日常配置经过这些坑我固定了一套排查工具链。首先是torch.cuda.memory_summary()每次显存异常时第一个调用。其次是py-spy dump用来抓取Python线程的堆栈定位卡死位置。然后是nvidia-smi -l 1实时监控显存和GPU利用率。最后是nsysNsight Systems用来分析CUDA内核的执行时间找出计算瓶颈。这套工具链覆盖了从Python层到CUDA层的所有环节。我建议你在开发环境就配好这些工具不要等到线上出问题才临时找。另外日志里要记录每个请求的输入长度、输出长度、显存变化这样出问题时能快速定位是哪个请求触发了异常。
返回列表