☰
字节码虚拟机+JIT:破解动态张量形状下的推理性能瓶颈
2026/10/3 11:24:21 网站建设 项目流程

做推理引擎的人大概都有过这种体验:模型在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做特化:batch=16、seq=32的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,例如batch=32且seq=32时,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值按模型复杂度自动调节,简单模型K=2,复杂模型K=8。稳定窗口过滤了偶发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在等CPU,CPU在编译,全体阻塞。

改成后台编译后问题立刻缓解:执行线程遇到未命中shape时,先走解释器路径垫住延迟,同时把编译任务丢到一个带优先级的后台队列,编译完成后查漏更新缓存。这里有个小技巧:后台编译队列按“预估收益=shape出现频率×单次执行时间”排序,收益低的任务可以延迟甚至丢弃,避免编译线程被低价值任务占满。

坑现象根因修复
编译风暴CPU占用高、首包延迟大偶发shape也触发编译稳定窗口+按频率提升
guard误伤偶发显存越界未检查stride/layout强guard补layout检查
同步编译偶发卡顿数秒编译占执行线程后台编译+优先级丢弃

5. 实测效果:一个动态batch Transformer的编译-执行拆解

5.1 测试条件与模型配置

为了讲清楚这套系统到底值不值,我搭了个动态batch的6层Transformer做基准,hidden_size=256,head=8,sequence固定为32,batch在1、4、16、64之间按随机游走变化。对比四组:纯PyTorch eager模式、TorchScript固定shape的静态编译、我的字节码VM+JIT、以及关闭JIT只跑解释器的VM。统一跑500次前向,前50次作为预热。

5.2 数据结果与解读

方案稳定态单次延迟热点shape切换后的首包延迟总编译耗时平均显存占用
PyTorch eager2.1 ms2.1 ms0 ms1.9 GB
TorchScript(shape固定)1.1 ms—5.3 s2.2 GB
纯VM解释器1.7 ms1.7 ms0 ms2.0 GB
字节码VM+JIT(本文方案)1.2 ms2.9 ms1.4 s2.1 GB

先别急着看数字,我挑几个值得玩味的点。

第一,纯解释器VM比Eager快0.4ms。原因不是指令执行快,而是省掉了Python侧大量动态分发和shape推断,说明哪怕不做JIT,VM这一层本身就有价值。第二,VM+JIT的稳定态延迟接近静态编译的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. 再往前走:动态时代的VM+JIT还能怎么进化

6.1 形状聚类的代码族复用

现在每个shape组合都对应一份特化代码,但真实模型里很多shape在“结构上”是相似的。比如batch=16、seq=36和batch=18、seq=36,它们的指令分布几乎完全一样,只是少数常量不同。我的想法是引入shape聚类,把这些结构相似组合归到一个“代码族”,共享编译产物,只在入口处把常量差异作为参数传入。这个方向能把编译次数再降一个量级。

6.2 与显存规划联动,减少cudaMalloc的隐形开销

特化代码最大的好处不只是快,而是可预测。当JIT提前知道某个shape组合会频繁出现,就意味着一批中间张量的大小和生命周期是已知的,可以用一个专用arena缓存这些张量的显存,避免每次调用都走一次显存分配器。我在实验里看到,配合arena缓存后,动态batch场景的显存分配次数下降了百分之八十。这个优化好玩的地方在于:它不是因为“代码更好”,而是因为“运行模式更可预测”。

6.3 给同行一个实在的建议

如果你也要做一个类似的系统,我的核心建议只有一条:先做对guard,再谈编译。Guard是整个JIT的地基,shape、layout、device、dtype这四类检查搞清楚了,后面的特化代码才敢放开手脚。另外一个经验是,性能剖析别一上来盯指令执行时间,先看dispatch次数和显存分配次数,这两个才是动态张量场景里最容易偷走性能的黑洞。

这套组合折腾下来,我最大的体感是:编译器真正的价值不在于把所有代码编得飞快,而在于精确判断什么时候该编译、什么时候不该编译。字节码VM提供结构稳定性,JIT提供形状适应性——两者合在一起,动态张量计算才真正变得既可以预测,又足够灵活。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询