资讯动态

字节码虚拟机+JIT:破解动态张量形状下的推理性能瓶颈

发布时间:2026/10/3 11:24:26 来源:尧图企业网站定制
做推理引擎的人大概都有过这种体验模型在GPU上跑得好好的一到线上batch不固定延迟就突然涨三五倍。这其实是动态张量计算的老问题。模型前向过程中张量形状随输入变化——NLP里的序列长度不定、推荐系统的batch随时变化、图神经网络的节点数每批都不一样。传统静态编译可以让固定形状的计算跑到飞起可一旦形状不规则编译器很多判断就全塌了。我当时的解决思路是把模型先表示成稳定的字节码再在运行时用实时编译JIT根据实际观察到的形状做热点特化这就是标题里“字节码虚拟机实时编译”要做的事。这篇文章写给两种人一是不满足于调包、想自己动手做推理引擎的工程同学二是想知道VM和JIT到底怎么配合解决动态形状问题的系统爱好者。我下面先讲清楚动态形状为什么难缠再讲字节码VM在我的方案里如何承担“稳定层”然后展开实时编译管线和那些写进代码后又被逼着改掉的设计。1. 动态形状带来的麻烦比大多数人以为的更深1.1 动态到底“动”在哪三个层面动态张量计算里的“动态”不只是一个词。拆开看有三种典型情况处理办法完全不同。第一种是未知维。编译时刻根本不知道某个维度会是多少只有在运行期第一次拿到输入时才确定。这种情况在深度学习框架里最常见比如模型接收的外部输入其batch size在设计时就是-1或None。第二种是可变长维。维度值是动态的但取值范围可控。NLP里被打包到同一个batch的句子如果不做padding统一长度那么seq_len就会在8到128之间随机波动推荐系统里每个用户交互的物品数量也长短不一。第三种是结构动态。不仅数值维度变化整个计算图的分支都可能改变。典型代表是控制流while循环的迭代次数、if条件的走向、动态shape下concat或split出来的张量集合大小都会影响后续指令序列。这三种动态有着完全不同的“伤害程度”。未知维还好处理因为它只是一个占位符一旦真值出现还能走静态路径可变长维最尴尬因为它会让任何基于shape的缓存以很高概率失效结构动态则直接挑战编译器的控制流分析能力。我的一个直观感受是很多人以为dynamic shape就是个“运行时才知道维度”的小问题但在真实系统中它其实是内存分配、算子选择、融合策略三件事的连环击穿。1.2 为什么静态编译器碰到动态形状就“破功”静态图编译器习惯在“编译期”就把一切都定下来。以XLA这类系统为例它会对输入shape做特化batch16、seq32的Transformer会被编译成一整套针对这个形状的优化代码从内存布局到kernel选择全都固定死。这样做的上限极高但代价是换一个shape就得重新编译一次。动态形状让静态编译的三大支柱同时失效第一根支柱是内存规划。静态编译可以用线性内存分配、arena等技术把所有中间张量的偏移量在编译期算好shape一旦动态这些偏移量就变成运行期变量要么走慢速的动态分配要么预分配一个上限很大的缓冲池浪费显存。第二根支柱是kernel选择。cuDNN里的卷积和矩阵乘实现最优先的算法往往依赖精确的shape、对齐方式和batch大小shape动态变化时要么每次重新做benchmark要么退回一个安全但未必最快的通用实现。第三根支柱是算子融合。融合的目的不只是少几次kernel launch更是为了把中间结果留在寄存器或shared memory里。而绝大多数融合算子要求输入输出shape完全静态哪怕中间多了一个torch.size()的读取都能让融合链条断掉。所以在我做系统设计时第一原则就是不要尝试在编译期解决所有动态问题要把动态当作“运行期的一种常态”来设计让VM和JIT去承接它。1.3 真实场景案例动态形状不是刁难是业务常态很多人有一个误解觉得动态形状只出现在“研究用的demo里”。恰恰相反生产系统里几乎没有静态形状。推荐系统最有代表性。线上请求的batch由并发量决定一般服务端会凑一批请求再推理batch从1到64随机分布特征侧的embedding lookup结果拼接后第二维长度也跟着变。在这个场景里哪怕模型权重完全固定输入和中间张量的形状每时每刻都在变。NLP场景也一样如果不做固定长度截断长短句混合的batch会让attention矩阵的形状完全不可预判。图神经网络更直接每个batch包含的图节点数、边数都不一样消息传递产生的张量shape随拓扑而变。在这些场景下系统要的不是“这一种shape跑到极致”而是“一批常见shape都能跑得快并且切换时不能卡顿”。这正是字节码虚拟机加实时编译能发挥价值的地方。2. 字节码VM在动态场景中的定位比AST解释器更稳比静态图更活2.1 为什么中间层是“字节码”而不是AST或纯机器码动态张量计算有两条极端的实现路线一条是纯解释执行比如Python里逐算子分发灵活但每条指令都有解释开销另一条是纯静态编译性能高但动态形状会带来重新编译风暴。字节码虚拟机站在两者中间但它不是简单“折中”而是把“表示”和“执行”解耦。先看为什么不用AST。AST的问题是节点带有强语法语义解释AST时递归深度不确定、类型信息不规整编译器要遍历树做分析和变换也麻烦。对运行时系统来说AST是一条“建模不彻底”的中间产物。而字节码是线性指令流每个op都对应明确的输入输出、属性和跳转关系天然适合做数据流分析、控制流分析和后续代码生成。再看为什么不直接生成机器码。动态shape下机器码的特化粒度很难把握特化太粗换shape就失效特化太细代码膨胀且缓存命中率低。字节码只在“结构层面”固定模型所有shape相关的决策都推迟到运行期由JIT按需特化。这样既能保证结构稳定又保留运行时灵活性。我当时设计的字节码分三类指令。第一类是tensor运算指令例如MATMUL dst, lhs, rhs带泛型属性和可选的shape标注第二类是控制流指令例如JumpIfFalse cond, target、LoopBegin、LoopEnd第三类是shape管理指令例如ShapeOf、ReshapeIf、ShapeAssert。这三类合在一起VM既能表达复杂模型又不会像AST那样把运行时和语法耦合在一起。2.2 栈式还是寄存器式这个选择直接影响编译期分析字节码虚拟机有个经典分歧栈式还是寄存器式。Python和Java的VM选了栈式Lua和Dalvik则偏向寄存器式。我在张量计算场景里最终选了寄存器式原因很实际。栈式指令的优点是字节码紧凑、生成器简单缺点是每条指令都要做栈的push和pop数据流关系隐含在栈序列里编译期要重建use-def链比较费劲。寄存器式指令流虽然每条指令都要带操作数索引字节码体积更大但它把数据流关系直接摊开了。对于JIT来说看到一个MATMUL %r2, %r0, %r1就能立刻知道张量的来源做常量传播、shape推断和kernel融合就简单很多。在支持动态shape方面我额外设计了一条叫ShapeGuard %r_dst, %r_src_shape, expected_pattern的指令。它在运行时把实际shape和期望模式做一次匹配匹配成功就继续失败就跳到重新编译路径。因为shape本身也是一个一等对象guard的输入是张量的shape对象而不是逐维比较的标量这样能把检查开销压到极低。2.3 VM运行时只做三件事分派、保shape、触发编译真实跑起来后我的VM运行时只干三件事。一是常规分发。把字节码指令分发到对应的kernel上这部分要尽量薄能直接调C实现就不要套多层虚函数。二是shape一致性维护。每个张量后面挂一个shape对象形状变化时由专门的指令更新而不是让每个算子各自维护一份。三是在遇到没见过shape的组合时触发JIT编译并把编译完成的特化函数注册回字节码指令对应的cache槽里。设计VM时最常犯的错误是让它承担太多职责。如果VM负责了内存管理、并发调度、错误恢复、和宿主语言交互那它就会变得越来越重丁点性能问题都说不清来源。我后来硬性规定VM只管指令流和shape其他都归JIT管。3. 实时编译管线的关键设计从字节码到特化机器码3.1 触发策略第一次遇见就编译那你就等着编译风暴吧字节码VM给出了一种很自然的JIT触发点某条指令第一次执行时它的cache槽是空的这时候可以选择编译。我一开始很天真把“第一次见就编译”当成默认策略。结果很惨模型里同时有10个动态shape维度组合起来有几百种预热阶段CPU忙到冒烟GPU反而在空转用户看到的不是变快而是明显卡顿。后面我把触发策略改成两级。第一级叫“记录”运行期只记录每条字节码指令见到的shape组合用哈希表按频率计数不触发任何编译。第二级叫“提升”当某个shape组合在滑动窗口内连续出现N次且它的累计执行时间超过阈值才把它从待编译集合提升到待编译队列。这个策略的本质是不为偶发shape付编译费只为稳定出现的shape付。3.2 特化编译到底特化了什么当某个shape组合被选中JIT要做的第一件事是从字节码片段构建数据流图。得益于寄存器式设计这一步几乎不用额外分析指令里的use-def链就是图。接下来针对实际shape做三件事。第一把shape相关的标量计算折叠成常量比如在编译期算出N*N、seq_len*head_dim这些固定值运行时不再重复计算。第二为固定shape挑选或生成专用kernel例如batch32且seq32时attention的softmax可以走向量化更激进的实现。第三把相邻的张量算子融合成一个大kernel减少显存往返和kernel launch次数。特化函数的入口名我习惯写成kernel_op_shape_hash放在一个全局的cache map里key是指令ID加shape组合的hash值。这听起来简单但真正写起来要注意hash函数得足够快不能在运行时浪费几十个纳秒。3.3 Guard体系不是越严越好特化代码运行的前提是“运行环境没变”而这个前提靠guard来维护。Guard是JIT系统里最容易做过头的地方。我见过有同行在guard里逐维度比较shape、比较dtype、比较内存连续标志、还比较设备ID一个guard执行出去几十纳秒没了反而比不特化还慢。Guard要分级。强guard针对高性能特化kernel检查shape、dtype、contiguity和device。弱guard针对只做轻量优化的路径仅检查rank和元素总数。选择哪种由编译时对kernel性能收益的预估决定收益大就上强guard收益小就老老实实用弱guard。另外是一个非常实用的优化多个张量如果共享同一个shape对象guard就不必逐个比较张量而可以比较shape对象的指针。真实模型里一个算子输出的shape经常被下游好几个算子共享指针比较比逐维整数比较快一个数量级。3.4 去特化与回退机制JIT系统还必须处理“特化失效”的情况这就是去特化deoptimization。例如guard失败意味着当前输入不满足特化假设此时要能快速回退到字节码解释器而不是在特化代码里报错。我的做法是两级执行模式。默认先跑解释器JIT编译完成后把特化函数挂到缓存槽里但解释器仍然保留。Guard失败的路径不是立即报错而是“降级到解释器”然后由JIT决定是否对新的shape组合重新编译。这里有一个被很多人忽略的点去特化必须轻量到“秒回”因为它可能发生在任何一次shape变更时。如果去特化还要处理栈展开、寄存器还原、对象状态回滚那模型一换batch就相当于付一次完整的上下文切换开销。我在设计字节码指令集时刻意把guard失败时的跳转目标放在附近让它只需要改变PC和cache槽索引不做状态序列化把开销压到接近一条分支指令。3.5 编译后端怎么选LLVM、MLIR还是自研小后端JIT后端选择上我提供三档方案给不同场景。第一档是直接生成C代码拼成字符串后调用底层的运行时编译器编译。优点是实现最简单调试直观缺点是编译延迟高不适合对首包延迟敏感的服务。第二档是灌给LLVM的ORC JIT。灵活性高能复用优化管线但工程复杂度陡增而且LLVM版本升级容易带崩整个项目。第三档是自研针对固定模式的轻量后端。对Transformer这类结构规整的模型算子模式其实有限可以手写少数几个kernel模板用shape哈希去模板实例化不需要完整的编译优化流程。我最终的生产版本是第二档和第三档混用外层模板特化处理高频小kernel内层LLVM处理大段融合代码。混用的原则很简单如果一个特化函数编译时间超过它预计能省下的执行时间就不该编译它。这个判断可以在触发阶段做能筛掉大量无意义的编译请求。4. 冲进实现阶段后我踩的三个大坑4.1 shape抖动导致的编译风暴真实线上数据不会像测试集那么友好。我第一个生产版本上线后发现batch从1到64之间来回抖动但频率分布很散没有哪个shape能连续出现几次。结果就是系统永远在“记录-提升-编译”之间打转编译线程跑满GPU利用率反而掉下来。我最后上了“shape稳定窗口”机制同一个shape组合必须在一个滑窗内出现至少K次这个K值按模型复杂度自动调节简单模型K2复杂模型K8。稳定窗口过滤了偶发shape让JIT只围绕真正的热点shape工作。上线后编译总量下降了百分之七十多而热点shape的命中率反而提高了。4.2 guard只查shape不查layout直接踩进显存陷阱我自己写过一段很蠢的代码guard只检查了tensor.shape和dtype没有确认内存layout。结果模型里一个view操作把shape改了但底层存储还是旧布局特化kernel按新的连续布局去访问显存直接越界。排查了整整两天才意识到shape相等不代表内存布局相等。同一组维度在不同stride下访存模式完全不同。我在guard里补上了stride检查和存储偏移一致性检查并把这类检查归入强guard只有高性能kernel才用。顺便说一句这类bug的隐蔽性极强普通测试很难触发因为它和PyTorch的view语义、显存分配器的复用行为都有关联。4.3 把编译放在用户线程上延迟就会失控最初版本的JIT编译直接跑在执行线程里。遇到新shape时VM在编译完成前会同步阻塞。这个设计在离线验证时没暴露问题到了线上有用户反馈“偶尔一下卡两三秒”百思不得其解。后来意识到编译线程抢占的是执行线程的时间片GPU在等CPUCPU在编译全体阻塞。改成后台编译后问题立刻缓解执行线程遇到未命中shape时先走解释器路径垫住延迟同时把编译任务丢到一个带优先级的后台队列编译完成后查漏更新缓存。这里有个小技巧后台编译队列按“预估收益shape出现频率×单次执行时间”排序收益低的任务可以延迟甚至丢弃避免编译线程被低价值任务占满。坑现象根因修复编译风暴CPU占用高、首包延迟大偶发shape也触发编译稳定窗口按频率提升guard误伤偶发显存越界未检查stride/layout强guard补layout检查同步编译偶发卡顿数秒编译占执行线程后台编译优先级丢弃5. 实测效果一个动态batch Transformer的编译-执行拆解5.1 测试条件与模型配置为了讲清楚这套系统到底值不值我搭了个动态batch的6层Transformer做基准hidden_size256head8sequence固定为32batch在1、4、16、64之间按随机游走变化。对比四组纯PyTorch eager模式、TorchScript固定shape的静态编译、我的字节码VMJIT、以及关闭JIT只跑解释器的VM。统一跑500次前向前50次作为预热。5.2 数据结果与解读方案稳定态单次延迟热点shape切换后的首包延迟总编译耗时平均显存占用PyTorch eager2.1 ms2.1 ms0 ms1.9 GBTorchScriptshape固定1.1 ms—5.3 s2.2 GB纯VM解释器1.7 ms1.7 ms0 ms2.0 GB字节码VMJIT本文方案1.2 ms2.9 ms1.4 s2.1 GB先别急着看数字我挑几个值得玩味的点。第一纯解释器VM比Eager快0.4ms。原因不是指令执行快而是省掉了Python侧大量动态分发和shape推断说明哪怕不做JITVM这一层本身就有价值。第二VMJIT的稳定态延迟接近静态编译的1.1ms只慢0.1ms而代价是总编译耗时只花了1.4秒远低于TorchScript重新编译一整遍的5.3秒。第三热点shape切换后的首包会到2.9ms因为要做一次guard失败、回退解释器、可能触发后台编译的流程。但接下来同一shape的第二次调用就能恢复到1.2ms这个恢复速度是静态编译做不到的。5.3 调优参数的经验值经过多轮压测我总结出三个最值得调的旋钮。第一个是guard强度分级开关。如果模型访存模式规整强guard的多检查开销完全可接受如果模型里大量用view、transpose建议把部分路径降级为弱guard宁可在极少数case下重复编译也不要频繁误伤。第二个是编译触发截断。我在JIT里设了“每指令最多缓存16个shape组合”的上限超过后走LRU淘汰。动态batch场景下这个上限控制在8到16最稳太小会让热点shape被挤掉太大又会让缓存占满内存。第三个是后台编译的并发度。实测单编译线程就够了并发超过2时收益趋近于零反而增加锁竞争和CPU上下文切换。6. 再往前走动态时代的VMJIT还能怎么进化6.1 形状聚类的代码族复用现在每个shape组合都对应一份特化代码但真实模型里很多shape在“结构上”是相似的。比如batch16、seq36和batch18、seq36它们的指令分布几乎完全一样只是少数常量不同。我的想法是引入shape聚类把这些结构相似组合归到一个“代码族”共享编译产物只在入口处把常量差异作为参数传入。这个方向能把编译次数再降一个量级。6.2 与显存规划联动减少cudaMalloc的隐形开销特化代码最大的好处不只是快而是可预测。当JIT提前知道某个shape组合会频繁出现就意味着一批中间张量的大小和生命周期是已知的可以用一个专用arena缓存这些张量的显存避免每次调用都走一次显存分配器。我在实验里看到配合arena缓存后动态batch场景的显存分配次数下降了百分之八十。这个优化好玩的地方在于它不是因为“代码更好”而是因为“运行模式更可预测”。6.3 给同行一个实在的建议如果你也要做一个类似的系统我的核心建议只有一条先做对guard再谈编译。Guard是整个JIT的地基shape、layout、device、dtype这四类检查搞清楚了后面的特化代码才敢放开手脚。另外一个经验是性能剖析别一上来盯指令执行时间先看dispatch次数和显存分配次数这两个才是动态张量场景里最容易偷走性能的黑洞。这套组合折腾下来我最大的体感是编译器真正的价值不在于把所有代码编得飞快而在于精确判断什么时候该编译、什么时候不该编译。字节码VM提供结构稳定性JIT提供形状适应性——两者合在一起动态张量计算才真正变得既可以预测又足够灵活。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑