资讯动态

动态张量计算:用字节码虚拟机与实时编译打破静态图困局

发布时间:2026/10/4 5:41:06 来源:尧图企业网站定制
过去半年我一直在折腾一个听起来有点偏门的方向给动态张量计算做一个带字节码虚拟机的运行时再在这个虚拟机之上叠加实时编译能力。起因非常朴素——业务里一堆长尾模型输入形状跨度极大从几十个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的batch32时原生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, axis1 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-attentionshape每一步都在变。对这种结构纯线性特化就失效了——你没法无限展开循环体。我目前的处理方式是循环存在但循环次数由运行时寄存器决定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 eager58.2基线无静态图 padding到51247.6基线 28%完整编译一次静态图 分桶8桶35.1基线 12%8次编译VM解释执行51.4基线 5%无VM JIT特化编译21.8基线 9%数百次编译这个结果基本验证了设计预期。JIT模式相比eager有2.7倍加速相比静态图分桶还有明显优势。原因在于分桶方案只能粗粒度匹配桶内形状仍然参差不齐kernel调用时要做runtime dispatchVM的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做AOTAhead-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算子链——从输入到输出——再开始做优化否则你会被性能需求绑架设计。这个顺序颠倒一次返工成本真的很高。

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

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

免费获取报价 →
↑