过去半年我一直在折腾一个听起来有点偏门的方向:给动态张量计算做一个带字节码虚拟机的运行时,再在这个虚拟机之上叠加实时编译能力。起因非常朴素——业务里一堆长尾模型输入形状跨度极大,从几十个token到上千个token都有,用PyTorch的动态图模式跑,GPU利用率经常只有两三成;换成静态图方案,形状一变编译就失效,被狠狠按回了解释执行。被逼到墙角之后,我决定从底层把这套链路重走一遍:不是去改造某个框架,而是自己写一个"能看见形状变化的虚拟机"。
这套系统做了大概五个月,目前已经在内部几个变长输入的推理服务上跑起来了。简单说,它做的事情是:把一个模型编译成指令序列,交给字节码虚拟机解释执行;同时虚拟机在运行时监测每个张量的实际形状,一旦形状稳定下来,就触发实时编译,针对当前这组具体形状生成优化过的执行路径。整个过程自动完成,不需要用户指定任何动态轴的边界值。
这篇文章我会把完整的设计思路和踩坑过程写出来,包括指令集长什么样、JIT策略怎么定、内存池怎么做、以及实测中那几次差点把缓存打爆的事故。适合正在做推理引擎、编译器后端,或者被动态形状恶心过的框架开发同学参考。
1. 动态张量计算的性能困局:静态图为什么拿它没办法
1.1 动态形状到底动了谁的蛋糕
先说清楚"动态张量计算"具体指什么。一个张量的形状如果在编译期无法完全确定,那么它就算动态张量。典型场景就三个:NLP变长序列、图神经网络、稀疏输入。以NLP为例,一个batch里每条样本的token数量不同,经过padding和mask处理之后,encoder部分的序列长度其实还是浮动的,更不用说decoder阶段逐个token生成时序列长度根本就是递增的。
这种动态形状对编译器和运行时都是考验。静态图框架做优化时有一个隐含前提:所有张量的shape在编译期已知。算子融合、内存规划、kernel选择、并行切分,全部依赖这个前提。shape一旦变成运行时变量,这些优化要么退化成保守策略,要么干脆失效。你没法在编译期决定"这个matmul到底调用哪个kernel",因为M、N、K的数值你根本不知道。
我在项目初期做过一个摸底实验:同样一个两层MLP,固定shape的batch=32时,原生PyTorch eager模式跑一次前向大约0.8ms;把batch改成不稳定浮动值后,同样的代码单次执行时间跳到2.1ms以上。这个差距并不全是kernel本身变慢,而是调度逻辑、缓存miss、kernel选择分支这些杂七杂八的开销全被放大了。
1.2 eager模式和静态图在这个问题上的各自局限
动态图框架(PyTorch eager、TensorFlow eager等)的灵活性毋庸置疑,逐算子解释执行,用户怎么写都行。但代价是每次算子调用都要穿过Python解释器、dispatcher、kernel选择三层。模型一深,这个调度开销占比就非常难看。更关键的是,动态图模式默认把每个算子当作独立事务来执行,算子之间完全不做融合,中间结果频繁写回显存,带宽就这么白花花地浪费掉了。
静态图编译的优化能力强,可一旦遇到动态shape,常见的做法就只剩两个:padding到最大长度,或者把动态轴分桶。Padding这条路的算力浪费非常赤裸,最大长度3000的序列padding后跑了10个token的短样本,计算量膨胀三百倍,GPU利用率再高也是烧钱。分桶稍微好一些,但桶的边界设置是门玄学,桶太粗浪费资源,桶太细则编译数量爆炸,而且实际负载波动大的时候桶的有效命中率低得让人心碎。
1.3 字节码虚拟机的定位:在灵活与高效之间再踩一条路
我当时想的是另一个方向:能不能让运行时"边跑边看边编译"?程序的指令序列是固定的,但执行路径可以跟着形状自适应。这就是字节码虚拟机能发挥作用的地方。
VM天然适合这种场景。它一方面比Python解释器硬核得多——不需要经过Python对象系统,操作数直接是张量元数据和设备指针;另一方面又比静态图灵活——指令集里可以设计专门的shape查询、shape分支指令,让运行时对形状变化作出响应。传统VM是"解释执行一段字节码",我们做的VM在此基础上加上了"运行时形状监视 + 热点触发编译"的能力,相当于给VM装了一双眼睛。
我的理解是:字节码虚拟机在这里的角色,是充当解释器与静态编译器之间的中间层。它够灵活,能容纳动态形状带来的运行时不确定性;它又够结构化,指令序列里的每个操作都是可分析、可改写、可特化的。
2. 跑在VM里的张量程序:指令集与IR设计的取舍
2.1 为什么没直接用现成的IR
立项初期我想过直接用LLVM IR或者TVM Relay,后来都否了。LLVM IR本质上是面向标量和数组的,它没有张量语义。你没法在LLVM IR里自然地表达"这是一个shape可通过运行时查询的三维张量"以及"对它的dim 1做reduce"。强行用LLVM表达代价也很高,每个张量都要套一层数组抽象,后续特化编译的代码改写工作量大到无法接受。
TVM Relay倒是张量级别的IR,但Relay有个预设假设:形状尽可能静态,动态shape支持只是兼容能力。我们恰好相反,核心场景就是形状高度不确定的模型,IR设计必须把"形状在运行才定"当成头等公民来对待。抄接近的东西不如从头设计,所以最终选择了一个寄存器式的、面向张量运算的字节码指令集。
2.2 寄存器式指令集的骨架
选择寄存器式而不是栈式,理由很实际:指令密度高、翻译到后端更少间接跳转、操作数可以直接对应到SSA值。栈式VM(比如JVM)每条指令都要从操作数栈顶弹入弹出,张量作为一等对象的时候复制栈顶指针的开销很惊人;寄存器式则每个操作数都是显式的虚拟寄存器号,JIT编译映射到物理寄存器也就一步之遥。
我们指令集分四组。第一组是数据搬运类,LOAD_TENSOR把输入张量加载进寄存器,STORE_TENSOR保存结果,MOVE_REG做寄存器间拷贝。第二组是计算类,MATMUL、ELEMENTWISE_ADD、CONV2D、REDUCE_SUM这些,每个指令都带输入寄存器列表和输出寄存器,以及一组静态属性——比如matmul的transpose标志、conv的strides和paddings。第三组是形状操作类,RESHAPE、TRANSPOSE、BROADCAST_TO,这些指令在早期设计上是有讲究的:它们只修改张量的元数据(shape、stride),不触发实际数据搬运,真正搬运由后续的计算指令触发。第四组是控制流类,BRANCH_ON_SHAPE、JUMP、LOOP_BEGIN、LOOP_END。
# 一个变长序列reduce的简化指令序列 # reg0 存放原始变长输入,shape 由运行时决定 LOAD_TENSOR r0, input QUERY_SHAPE r1, r0.dim(1) # 把序列长度读到r1 BROADCAST_TO r2, r1, [1] # 把长度作为标量张量 MATMUL r3, r0, weights # 线性变换,shape 仍包含动态轴 RESHAPE r4, r3, [batch, -1] # 动态轴用 -1 占位 REDUCE_SUM r5, r4, axis=1 STORE_TENSOR output, r5这段指令看起来简单,但每一条的shape信息在编译指令序列时都是未知的。MATMUL指令的M不会写死在字节码里,而是运行时从r0的元数据里读取。这就是它与静态图IR最本质的区别。
2.3 张量作为一等公民:静态已知与运行时已知的边界
设计指令集时最重要的决定是:哪些信息编码进静态指令字段,哪些留在运行时元数据里。
我们的原则是"轴的粒度"上做切分。一个张量的三个维度,如果axis 0的size是编译期确定的(比如batch固定是1),那就把静态值直接写进指令的操作数属性里;如果axis 1的size是动态的,就通过QUERY_SHAPE指令在运行时拿到一个"尺寸寄存器",后续指令通过引用这个寄存器来获得动态信息。
这个设计对JIT非常重要。实时编译触发时,编译器扫描指令序列,凡是操作数里出现尺寸寄存器的地方,都用运行时实际拿到的值去替换。能静态化的全部静态化,剩下来不及静态化的,才保留运行时查询。整个特化过程有明确的方向——把每个张量的完全shape定下来,能定多少定多少。
2.4 我一定不会重做的几个设计决定
第一,不要在产品初版搞复杂类型系统。我一开始在指令集里设计了一套带layout信息的张量类型,结果前端翻译模型的时候被layout推断折腾得死去活来。第二,不要把优化策略写死在指令集里。比如融合、内存复用这些应该是JIT编译阶段的事,让指令集保持纯粹描述计算。第三,控制流不要一开始就做函数调用栈,动态循环已经够复杂了,先把扁平的基本块和跳转做扎实。
3. 实时编译的两条腿:解释执行兜底 + 形状特化编译提速
3.1 先解释后编译:VM执行模型的选择
这个VM的执行模型一开始就定成"两层"。第一层是纯解释执行——取指令、解码、根据shape元数据分派到对应后端kernel;第二层是实时编译——运行一段时间后,VM里的形状监视器发现某些输入的shape签名持续稳定,就把整段字节码针对这组具体shape做特化编译,生成内部执行上下文。
为什么不能直接先编译?因为动态shape的问题就是编译期信息不足。你连M和K都拿不到,怎么选kernel?所以必须有一段解释执行的过程来"观察"输入形状。对于很多推理服务,输入shape的分布其实非常集中,比如线上流量大部分样本长度集中在某个区间,观察几十个batch之后就能锁定常见形状。这个观察阈值是VM一个核心参数min_compile_iters,我们默认给50——也就是看到同一个shape签名连续出现50次后才触发编译。
3.2 形状特化编译的核心机制
触发编译后到底发生什么?我把这个过程拆成三步。
第一步是shape签名捕获。VM从当前执行流里取出所有输入张量的shape和dtype,组合成一个签名key,比如[fp32(16,128,768), i64(16), fp32(768,768)]。这个key就对应了一组完整的静态形状信息。
第二步是静态化重写。编译器遍历这段字节码,把所有依赖shape查询指令的寄存器,替换成第一步拿到的具体数值。QUERY_SHAPE指令被消除,MATMUL指令的M和K被填上具体数字。重写后的指令序列变成一个完全静态化的版本。
第三步是后端生成。针对静态化指令序列,我们调用底层算子库(BLAS、cuDNN或oneDNN)的具体kernel做绑定,消除这两层间接调度:原本解释循环里每次都要做的opcode分支,现在直接是一串连续的kernel调用;原本运行时才做的kernel选择,编译期就已经定好。
# 解释执行时每次循环都要做的分派 while (pc < code_len) { opcode = code[pc]; switch (opcode) { case OP_MATMUL: // 运行时读取shape,选kernel,然后执行 runtime_kernel_dispatch(regs, code + 1); break; ... } pc += insn_len; } # 特化编译后生成的逻辑 void compiled_entry_16_128_768(float* input, float* out) { // 所有shape都是常量:16, 128, 768 sgemm_16x128x768(input, weights, tmp); // 直接调用静态shape的kernel elementwise_add(tmp, bias, out); // 不再有opcode分派 }这段伪代码展示了这两种执行方式的本质区别。解释执行是"每次都要做决定",特化编译是"决定做一次,之后无脑执行"。对于一条100个算子的计算链路,解释执行至少有上千次分派判断,特化编译把这些全部消掉了。
3.3 编译缓存策略:shape key别做太细
缓存策略是JIT系统最容易翻车的地方。我第一版是把完整shape签名当key,精确匹配才复用编译产物。结果线上出现一个事故:模型输入长度在每个batch都微小浮动,比如[127, 129, 125, 128, 130],每个长度都触发一次编译,编译缓存里攒下几百个几乎一模一样的版本,不仅浪费编译时间,缓存内存也暴涨。这就是"shape抖动导致缓存爆炸"。
后来改了策略。第一步是对动态轴做桶化——把连续变化的值映射到16的倍数桶上,比如长度127和128都归到128的桶。第二步是限制编译总量,用LRU淘汰,最多保留32个特化版本。第三步是设置编译熔断:如果最近2分钟内新shape signature出现的频次超过阈值,就暂停编译,回到解释执行模式,等shape分布稳定了再恢复。这套"分桶+淘汰+熔断"的组合下来,缓存稳定性好了非常多。
3.4 控制流怎么办:动态循环和shape分支
动态张量计算里绕不开动态控制流。比如NLP decoder的逐步生成循环,循环次数取决于已生成序列的结束符位置,循环体里还带着self-attention,shape每一步都在变。对这种结构,纯线性特化就失效了——你没法无限展开循环体。
我目前的处理方式是:循环存在但循环次数由运行时寄存器决定,JIT编译时把循环体做特化编译(内部的shape全部静态化),循环次数保留为一个运行时的整数寄存器。解释器负责做循环控制流和shape元数据更新,JIT后的循环体内核则用最紧凑的kernel执行。这个混合模式在实测中能在"循环控制开销"和"kernel执行效率"之间取得不错平衡。
另外还支持一种BRANCH_ON_SHAPE指令,按照shape满足的条件跳转到两个不同的执行路径。比如长度小于100走EfficientAttention分支,大于100走标准Attention分支。JIT编译时如果发现某个分支在观察期内从没被走到过,可以选择不编译那个分支,省下一半编译时间。
4. 内存与调度:动态张量的另一座大山
4.1 动态shape带来的生命周期碎片问题
形状一变,中间结果的大小就跟着变。静态图框架可以在编译期把所有中间张量的shape算好,一次性规划出内存池,整个模型执行期间几乎不发生二次分配。但动态shape做不到——当前batch的Q矩阵可能比上一个batch大两倍,中间buffer必须动态扩容。
这也带来一个非常实际的问题:在高并发推理服务里,每个请求都是独立的模型实例,每个实例都有自己的动态中间张量。如果全部走系统的设备内存分配器,一旦并发上来,malloc/free潮汐式交替,分配器内部锁竞争和碎片率能把性能拉低一大截。
4.2 按shape分桶的专用分配器
解决方案是给虚拟机单独做一个设备内存池。这个分配器的逻辑非常直白:以(dtype, element_count)做key,把已经释放的块存放在一个哈希表的空闲桶里;申请的时候先查桶,有合适的块直接复用,没有再向底层分配器申请新块。
这个方案为动态shape场景带来了很大的好处——虽然无法像静态图那样精确规划每个buffer的地址,但至少每个常见shape都有了一块可复用的"专用空间"。举个例子,一个变长batch的[batch, seq, 768]中间张量,就算seq在80到150之间浮动,映射到桶上是有限的几种size,下一轮请求大概率能命中空闲桶。
4.3 调度模型:动态DAG的依赖感知
VM里每个算子不只是一个函数调用,而是一个带依赖关系的task。运行时维护一张轻量DAG:一个算子要执行,先要等它依赖的所有张量ready。算子调度器负责按DAG拓扑顺序把ready的算子发给后端线程池。
静态图框架可以提前做全图调度规划,把算子切到stream上异步流水。动态shape做不到完全提前规划,因为后续算子的shape依赖前序算子的实际输出shape。所以我选择的是分层调度:指令级按依赖就绪度实时调度,每个后端设备内部再做异步流水。每一层的粒度不同,但都是动态调度的,而不是静态规划的。
经验补充:后端线程池的并发度不要直接拉满。我遇到过同步原语竞争比kernel执行还贵的情况。对GPU来说,线程池最大并发度设定在4-6效果最好;对CPU推理,绑定到物理核数并关超线程会更稳定。你如果直接照搬ThreadPool的默认配置,性能会很难看。
5. 实测数据与踩坑记录
5.1 一组小基准:变长输入的Transformer encoder
我用一个12层Transformer encoder做测试,输入是长度浮动在64到512之间的随机token序列,batch固定为16。对比三条路径:PyTorch eager、静态图对最大长度512做padding、我们这套VM的JIT模式。
| 方案 | 单batch平均延迟(ms) | 额外内存占用 | 编译/预热时间 |
|---|---|---|---|
| PyTorch eager | 58.2 | 基线 | 无 |
| 静态图 + padding到512 | 47.6 | 基线 + 28% | 完整编译一次 |
| 静态图 + 分桶(8桶) | 35.1 | 基线 + 12% | 8次编译 |
| VM解释执行 | 51.4 | 基线 + 5% | 无 |
| VM JIT特化编译 | 21.8 | 基线 + 9% | 数百次编译 |
这个结果基本验证了设计预期。JIT模式相比eager有2.7倍加速,相比静态图分桶还有明显优势。原因在于分桶方案只能粗粒度匹配,桶内形状仍然参差不齐,kernel调用时要做runtime dispatch;VM的JIT是针对精确shape做了特化,kernel选择和调用路径都是直给的。
5.2 踩坑一:shape抖动导致的编译风暴
这是我在第3.3节提过的那个事故,值得展开说。上线第一天,某个模型输入长度不是聚集在几个固定值,而是均匀分布在30到280之间,几十万个batch里shape签名几乎每个都不重样。本地测试好好的,上线后就崩了——VM疯狂编译,CPU被打满,原本15ms的延迟飙到240ms。
事后复盘,当时的实现致命缺陷有两条:一是编译缓存无分桶无上限,二是触发条件过于宽松(min_compile_iters设成了10)。修法就是前面说的三件套:动态轴按16桶化、LRU最多保留32个特化版本、观测到高频新shape时熔断编译。另外还把min_compile_iters调回50。这个体验让我记了一个教训:JIT系统的所有参数默认值都要在实际线上流量分布下验证,本地基准测试的流量形态说明不了任何问题。
5.3 踩坑二:冷启动期的前几个批次的昂贵编译
JIT特化的代价是"第一个遇到这个shape的时候很慢"。我们观察到,冷启动阶段最严重时前20个batch的P99延迟是稳定期5倍以上。因为每个新shape进来都要触发一次全链路编译。
缓解思路是"预编译常见shape + 解释执行预热"。具体做法:模型部署时跑一个profile脚本,统计历史流量里shape签名的Top-K分布,把Top-8做AOT(Ahead-Of-Time,提前编译)直接缓存成特化版本。线上遇到这些shape直接命中,没有编译开销。遇到未覆盖的shape,退化为解释执行,同时异步触发JIT编译,下一次遇到同样shape就能用上。AOT和运行时JIT的产物完全共用一套缓存机制,实现成本不高,收益非常显著。
5.4 踩坑三:动态循环里shape推断的迭代收敛问题
JIT在动态循环里做shape特化时,最麻烦的是循环体内部还可能再做RESHAPE。比如循环体里有一个reshape把[batch, seq, 2*head_dim]拆成[batch, seq, head, 2, head_dim]这种操作。JIT编译时需要一个shape推断过程,算出reshape前后shape的映射关系。
但动态循环里,某个寄存器的shape可能依赖上一个循环迭代的输出shape。如果没有收敛控制,shape推断会一直迭代下去,严重时直接进入死循环。我们的解法是给shape推断加定点迭代的机制:shape值的变化范围必须单调收缩,如果某次推断导致shape增长,立即终止特化编译,放弃这轮JIT,回退到解释执行。虽然偶尔会丢掉一些优化机会,但稳定性优先,这比强行优化然后产出错误shape安全得多。
写在最后的一些心得
回头来看,给动态张量计算写字节码VM和实时编译,本质上是在"灵活"和"高效"之间找一个动态平衡点。静态图选择的是在编译期锁定一切,我们选择的是在运行时观察、然后针对观察到的规律做特化。这个思路不只适用于张量计算,任何存在运行时特征的程序都值得想一想:能不能先用解释器兜底,再用JIT追着热点打。
如果你也准备做类似的东西,我最后提醒三件小事:第一,shape key的粒度一定要经过线上数据验证,别想当然;第二,JIT编译的频率和时机做保守一点,恢复解释执行的能力永远比强制编译重要;第三,不管设计多漂亮的指令集,先跑通一条简单的动态shape算子链——从输入到输出——再开始做优化,否则你会被性能需求绑架设计。这个顺序颠倒一次,返工成本真的很高。