ARTICLE DETAIL

资讯详情

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

LLM推理优化实录:一行NumPy让采样提速千倍

LLM推理优化实录:一行NumPy让采样提速千倍 上个月我在调一套基于 sglang 的 LLM 推理服务遇到一个非常典型的“模型快、服务慢”问题。GPU 出 token 的速度一点没变但整条请求链路被拖到离谱单次生成从 400ms 飙到 10 秒。排查了两天才确认瓶颈不在显存、不在 batch、也不在框架配置而在推理之后的一小段 Python 后处理代码。更讽刺的是罪魁祸首是一个我自己写的、看起来完全正确的采样函数。把它核心逻辑换成一行 NumPy 调用后单次采样从 800 多毫秒掉到 0.2 毫秒左右性能提升数千倍。这篇实录就把定位、拆解、替换和验证的完整过程复盘一遍给做 LLM 推理服务优化、自研采样逻辑、或者魔改生成流程的朋友做个参考。1. 症状初现GPU 只花 200msPython 后处理却卡了足足 7 秒1.1 实验背景魔改采样器后的诡异延迟当时我在做生成阶段的自定义温度衰减需要在每个解码步拿到 logits按业务规则调整分布后再采样。sglang 默认的采样器不好扩展我就把服务拆成两层框架负责预填充和增量推理logits 返回给 Python 侧我自己写了采样逻辑做后处理。结果一压测就发现问题。模型侧 decode 一个 token 大约 40ms 级别按正常理解生成 100 个 token 也就 4 秒多。可实际单次请求动辄 8-10 秒而且主要时间消耗非常不均匀有时候前几个 token 很快突然某一个 token 要卡一两秒。这已经完全超出“模型变慢”能解释的范围因为模型推理速度是稳定的波动只能来自后处理链路。我第一反应怀疑是显存碎片或者 batch 调度问题调了半天没效果。后来在服务日志里加上分段时间戳才发现真正的问题有一个请求里GPU 侧总共只花了 200 毫秒而 Python 侧某个隐藏的函数吃了整整 7 秒。这个函数不是主路径上的显眼角色平时日志也不会单独记录它所以之前完全没暴露。1.2 cProfile 出手罪魁祸首是一个“看起来很正常”的 sample 函数定位这种性能黑洞最直接的工具就是 cProfile。用法很简单python -m cProfile -s cumulative sglang_wrapper.py-s cumulative表示按累计耗时排序跑完一轮请求后输出会告诉你每个函数到底吃了多少时间。我那个 case 的结果摘出来是这样ncalls tottime percall cumtime percall filename:lineno(function) ... 1000 856.2 0.856 856.2 0.856 sample_with_temperature整整 97% 的时间都花在sample_with_temperature这个函数里。问题是它是我照着教科书写的temperature 缩放、softmax、累积概率分布、随机抽样每一步逻辑都挑不出毛病。函数也不长十几行而已。就这么一个“模范生”成了整个服务的性能灾难。后来复盘为什么没早点发现。因为服务对外只打 tokens/s这个指标主要由 GPU 解码速度决定。采样函数的时间被算进了“解码之外的其他开销”除非把每一段都掐表否则根本看不出来。日志掩盖了真实分布这是分布式服务性能排查里最常见的盲区。1.3 为什么这种病很难一眼看出来逻辑正确的代码性能突然爆炸这类问题在代码评审里几乎一定会漏掉。平时 review 看的是功能边界、异常处理、内存释放很少会有人盯着一个for i in range(len(logits))去算它的总迭代次数。而且后处理链路的调用栈通常很深前面有 HTTP 路由、鉴权、tokenize、batch 组装后面才轮到它。人的注意力天然会放在入口和出口中间那段“看起来没毛病”的代码反而最容易成为盲区。这次之后我养成一个习惯任何涉及数据量大的后处理代码第一版跑通后都会顺手看一眼复杂度。大模型的词表动辄几万到十几万哪怕只是简单遍历一遍Python 循环的消耗都远超直觉。2. 病根分析用 Python 循环遍历 15 万词表每一步都在为动态类型买单2.1 这段代码实际在做什么先看原始代码一个典型的 temperature 采样函数import math import random def sample_with_temperature(logits, temperature): logits [x / temperature for x in logits] max_logit max(logits) exps [math.exp(x - max_logit) for x in logits] total sum(exps) r random.random() * total cum 0.0 for i, e in enumerate(exps): cum e if cum r: return i return len(exps) - 1逻辑上它做了三件事把 logits 按 temperature 缩放减去最大值防止exp溢出算 softmax 分母后做累积概率抽样。无论从数值稳定性还是抽样正确性上这个实现都没有问题。很多深度学习教程给出的采样代码和这个几乎一模一样。问题不在逻辑而在执行方式。它把“对 15 万个元素做向量运算”这件事硬生生拆成了 15 万次独立的 Python 循环迭代。每一步迭代都有完整 Python 解释器开销这个开销比 C 循环同一个操作要贵几个数量级。2.2 十几万次迭代的隐性成本对象装箱、GIL、内存分配主流大模型的词表大小LLaMA 系列大约是 32KQwen 系列常见的是 151936GPT 级别的大约 100K 上下。越大的词表这个函数的痛感越明显。每执行一次logits[i] / temperaturePython 都要做这些事情从列表里取出一个 PyFloatObject把浮点数装进新的 PyFloatObject返回值再次装箱引用计数增减。math.exp(x - max_logit)虽然最终调用 C 库但每次调用需要一次 Python 到 C 的上下文切换。累积抽样循环里的cum e同样要经历装箱和比较。也就是说一次采样要执行大约 15 万轮这样的操作每轮还都是多个 Python 级步骤的组合。我实测下来151936 词表下纯 Python 版本的采样函数单次耗时在 700ms 到 1.2s 之间波动。这还只是单线程的结果。如果服务开了多线程GIL 还会让这个数字雪上加霜线程越多竞争越严重采样函数越慢。对比一下 C 语言里同样的事一个 15 万长度的 float 数组除一个标量、调一次 exp、做一次累加编译器直接向量化几十微秒就能跑完。Google 的基准测试数据里Python 循环和 NumPy 向量化之间的性能差距通常在 100 到 1000 倍之间词表越大差距越夸张。2.3 为什么 NumPy 向量化能快三个数量级NumPy 的底层是 C 和 Fortran 数组所有逐元素运算都在连续的 C 内存空间里完成。logits / temperature不是一个一个地算而是一次 C 循环遍历整个数组中间没有任何 Python 对象分配。整个 15 万词的除法在 C 层面只是一层 for 循环几毫秒甚至更短。打个比方纯 Python 的做法是快递员一个人搬 15 万件包裹每搬一件都要弯腰、记账、回头确认NumPy 的做法是上一条传送带15 万件包裹一次性过去。搬运总量一样但中间的管理成本完全不同。还有一个非常关键的点NumPy 数组在内存里是连续存储的CPU 缓存的命中率远高于 Python listlist 里存的是指针数组真正的 float 对象散落在堆上。对大数组来说这个 cache 友好性的差距本身就有几倍。所以向量化不是“优化了一点”而是整个执行模型都变了。2.4 Gumbel-max 采样一行代码完成 temperature 采样的数学原理问题在于怎么把“按概率采样”也向量化。最朴素的想法是用np.random.choice(vocab_size, psoftmax(logits/T))这样确实把循环消灭了但计算 softmax 仍然需要算完整个指数、做完归一化仍然有额外开销。更快的方式是 Gumbel-max 技巧。它的结论很简洁给每个 logits 加上一个服从 Gumbel(0,1) 分布的随机噪声然后取最大值对应的索引结果等价于按 softmax(logits/T) 的概率做抽样。标准写法是先对 temperature 缩放后的 logits 加上独立同分布的 Gumbel 噪声再 argmax。数学形式是import numpy as np noise np.random.gumbel(sizelogits.shape) next_token int(np.argmax(logits / temperature noise))为什么可行因为 Gumbel 分布有一种“最大稳定性性质”一组带噪声的得分中最大得分对应的索引服从 softmax 给出的概率分布。也就是说argmax(logits / T Gumbel_noise)和“先算 softmax 再按概率抽签”在分布意义上是完全等价的而且省掉了 softmax 和累积抽样的全部中间步骤。这个技巧还有一个隐藏优势argmax对整体加减常数不敏感所以不需要像原始代码那样做max_logit防溢出处理。数值稳定性天生就更好。3. 那一行代码替换后的实测表现与细节处理3.1 替换前与替换后的代码对比替换后的完整函数长这样import numpy as np def sample_with_temperature_fast(logits, temperature): noise np.random.gumbel(sizelogits.shape) return int(np.argmax(logits / temperature noise))核心逻辑就一行np.argmax(logits / temperature noise)。原始版本 18 行替换版本 4 行。如果想把噪声生成也内联进去甚至可以压成一行next_token int(np.argmax(logits / T np.random.gumbel(sizelogits.shape)))实际项目里我会拆成两行可读性好一些也方便固定随机种子。如果你手里的 logits 还是 torch.Tensor先转成 numpy 数组再走这个函数即可logits_np logits.cpu().numpy()3.2 性能实测数据附表格测试环境Python 3.10 NumPy 1.2632 vCPU 的云主机词表大小 151936每个版本跑 500 次采样取中位数。结果如下实现方式单次采样耗时相对原始版提升纯 Python 循环含累积抽样约 865ms1xNumPy softmax np.random.choice约 0.42ms约 2000xGumbel-max 一行式约 0.18ms约 4800x注意这里的 2000x、4800x 是实测值不同机器、不同词表大小会有波动。标题里保守写 1000 倍是因为这个数字在大多数场景下都能稳定复现实际上往往更高。我特别说明一下为什么没直接推荐np.random.choice方案。虽然它也向量化了但它要先计算 softmax 得到完整概率数组再做一次 C 层面的抽样中间会多一次全数组的指数计算和归一化。Gumbel-max 直接跳过 softmax只做一次加减和一次 argmax省掉的时间在 2 到 3 倍左右。对于每 decode 一步都要调用的后处理函数这点差距相当可观。3.3 精度与随机性验证分布没有变种子仍然可复现性能提升这么大第一反应肯定是怀疑“抽样分布会不会变了”。我专门做了验证固定一组 logits两个版本各采样 10000 次统计每个 token 被抽中的频率。用 softmax 概率作为基准高频 token 频次的相对误差都在 0.5% 以内低频 token 因为抽样次数少会有正常的统计波动但整体分布完全一致。随机种子方面有个容易踩的坑Python 标准库的random.seed和 NumPy 的np.random.seed是两套独立的随机状态机。替换代码后原来依赖random.seed做可复现实验的话必须额外设置np.random.seed(some_seed)否则每次采样序列还是可复现但和旧实现的序列对不上。这一点对做实验、跑评测的时候尤其重要别等结果对不上才想起来。还有一个细节np.random.gumbel的采样很便宜但在超大批次、极高并发下频繁调用也会产生一定开销。如果你的采样频率极高可以预生成一批 Gumbel 噪声数组循环使用。不过绝大多数场景下没必要直接调用就行。3.4 同一思路外推top-k 和 top-p 过滤的向量化写法很多项目除了 temperature 采样还要做 top-k 和 top-p 过滤。这类操作同样经常被人写成 Python 循环完全没必要。top-k 过滤的向量化非常直接k 50 idx np.argpartition(logits, -k)[-k:] mask np.ones_like(logits, dtypebool) mask[idx] False logits[mask] -np.infnp.argpartition不是完全排序只保证第 k 大的元素在正确位置复杂度接近 O(n)比完整排序的 O(n log n) 快不少。top-p 过滤则可以用排序加累积求和sorted_logits np.sort(logits)[::-1] cumprobs np.cumsum(softmax(sorted_logits)) cutoff np.searchsorted(cumprobs, p) logits[sorted_logits sorted_logits[cutoff]] -np.infnp.searchsorted一次调用完成“找累积概率达到 p 的位置”比循环判断快几个数量级。这些操作组合起来整个后处理链路都保持在微秒到亚毫秒级别和原来秒级的体验完全不是一个量级。4. 第二个千倍案例处理 LLM 返回的非法 JSON 时正则灾难性回溯4.1 场景一段正则把后处理拖垮采样性能修完之后我又用同样的思路审查了整条后处理链路果然发现第二个雷JSON 提取与修复。LLM 返回的文本经常不是合法 JSON多一个引号、少一个括号、或者在外面包了 markdown 代码块都是日常操作。项目里写了一个基于正则的提取器用来从模型回复里抠出 JSON 对象。压测时发现正常情况下这个提取器只要 0.1 毫秒左右但只要模型输出里出现明显格式错误耗时立刻暴涨到 8 秒直接触发服务超时。提取器的核心正则类似这样import re JSON_PATTERN re.compile( r(\\.|[^\\])*\s*:\s* r(\[[^\[\]]*\]|\{[^{}]*\}|(\\.|[^\\])*) )它试图匹配“键 冒号 值”的结构值可以是数组、对象或字符串。看起来考虑得挺周全但它掩盖了一个致命问题嵌套结构一多正则引擎的回溯路径会指数爆炸。4.2 根因灾难性回溯的数学解释正则引擎在匹配失败时会回溯到之前的分支点尝试其他路径。如果正则里存在嵌套的可选分支和重复组在某些输入下尝试的路径数量会呈指数级增长。最经典的例子是^(a|a)$匹配aaaaaaaaaaaaaaaa!。当最终遇到!导致匹配失败时引擎要把前面 16 个a按各种方式切分成若干个(a|a)组合切割方式有 2^15 种然后逐一尝试。字符串越长组合数量爆炸式增长。这就是“灾难性回溯”。我那个 JSON 正则的问题类似它有(\\.|[^\\])*、\[[^\[\]]*\]、\{[^{}]*\}这样多个可重复的嵌套选择分支。当模型输出在某个引号处不闭合或者括号层级错位时引擎会尝试所有可能的划分方式回溯数量随着文本长度迅速飙到天文数字。600 个字符的响应已经足以让回溯时间长到不可接受。4.3 一行代码修复换掉那个嵌套贪婪正则修复思路不是“优化正则”而是彻底放弃用正则去理解嵌套结构。JSON 本身的语法是有递归性的正则不是处理这种结构的正确工具除非你用专门的递归下降解析器。我换成了非常朴素的方案先找到文本中第一个{和最后一个}截取中间部分作为最外层 JSON 候选再交给json.loads如果解析失败再用json_repair库做修复import json from json_repair import repair_json_string s text[text.find({): text.rfind(}) 1] try: data json.loads(s) except json.JSONDecodeError: data json.loads(repair_json_string(s))json.loads是 C 实现的确定性解析器复杂度线性json_repair是专门的错误容忍解析器处理常见格式问题比正则可靠得多。替换后即使是严重损坏的 JSON 响应整个提取和修复流程也从 8 秒降到了 1 毫秒以内。这又是一次“一行代码级别”的改动性能差距上千倍。4.4 复盘哪些“后处理”代码最容易埋雷经历这两次优化后我总结出 LLM 推理后处理链路里最三类容易性能爆炸的代码第一类是用 Python 循环遍历大词表或长序列的逻辑典型就是采样函数、argmax、过滤、归一化。这类问题通过向量化解决性能提升往往极其夸张。第二类是用正则处理嵌套结构典型就是 JSON 提取、代码块解析、括号匹配。正则在“完全匹配”的场景下很好用一旦输入来自一个有概率出错的模型灾难性回溯就会找上门。建议优先使用专门的解析器让正则只负责“找出候选片段”结构解析交给确定性工具。第三类是逐字符扫描的字符串操作比如 BOM 清理、编码检测、不可见字符过滤。这些操作看着人畜无害但字符串长度一旦到几十上百 KBPython 逐字符循环就会变成隐藏的秒级耗时。优先用str.replace、字节串的bytes.translate或正则的预编译模式。排查时我给自己的清单就三条有没有循环循环多少次有没有处理异常输入时的复杂度陷阱直接把这三条过完后处理链路的性能大头基本都能揪出来。最后分享一个小习惯。这轮优化之后我会给后处理链路的所有关键函数都单独加性能日志采样、解析、格式化各一段便于线上直接看到耗时分布。性能优化这件事最重要的不是“快”而是先知道时间到底花在哪。养成 profile 的习惯往往会在你最想不到的地方遇到一行代码改变整个服务吞吐的时刻。
返回列表