资讯动态

动态shape下的张量编译:字节码VM如何用实时JIT突破性能瓶颈

发布时间:2026/10/5 1:03:30 来源:尧图企业网站定制
1. 这个场景我为什么觉得值得做1.1 动态张量到底难在哪做编译器或者运行时的人应该都体会过这种憋屈模型在训练和推理时跑得好好的一旦碰到变长的序列、树结构的输入、图神经网络里不规则的邻居数量静态图那一套就立刻露馅。画面很常见——batch里每条样本长度不一样你只能padding到最长算力浪费一大半或者干脆放弃图编译退回收发器模式每一层都走Python调度性能断崖式下跌。这里说的“动态张量计算”不是指那种靠向量化指令加速的普通张量运算而是指张量的维度、形状、甚至是控制流走向在执行前无法静态确定。比如Transformer里decode阶段的cache增长、GNN里每个节点的邻居数不同、推荐系统里每个用户的交互序列长短不一这些都属于典型的动态shape场景。它们共同的特性是编译时不知道运行时才知道。传统编译器有个习惯就是喜欢把所有信息在编译期定死。Shape定死了循环上下界定死了寄存器分配、向量化策略、访存模式全都好办。可动态shape一进来这一切全得推倒重来。你不能把循环写成固定次数不能假设每个batch的layout都一致连中间buffer的尺寸都没法提前计算。静态编译器面对这种输入基本上只有两条路要么疯狂产生guard和fallback导致代码膨胀得离谱要么干脆放弃治疗把场景退给解释器。两条路都谈不上优雅。所以我们当时的目标很明确搞一套运行时系统既能保留JIT编译的性能收益又能优雅地应对“运行时才知道shape”的现实。这就需要一个不把shape当唯一依赖的中间表示也需要一套能在几微秒到几十微秒内完成编译决策的实时编译机制。简单说就是让编译器给运行时打工而不是反过来。1.2 静态编译器为什么不香了大家一提到高性能张量计算第一反应就是上XLA或者Glow那套静态图编译思路。静态编译确实猛把所有op都lower成LLVM IR或者底层kernel然后做跨op融合访存能省则省。但静态编译有它天然的天花板图输入必须是完整且固定的或者至少是“符号化可推导”的。问题在于现实中的动态shape根本没法完全符号化。你可以在IR里用符号变量代表一个未知维度但一旦这个符号变量出现在循环边界、reshape的目标形状、切片索引里推理就变得极其复杂。举个例子x.reshape(seq_len, -1)这种操作静态编译器看到-1就只有干瞪眼while i x.size(0): ...这种循环即便形状编译器能处理生成的代码也要为所有可能的x.size(0)准备一套逻辑最终结果要么是代码爆炸要么是性能回归到解释器水平。更核心的问题在于静态编译把“编译”和“执行”两个阶段切得太开。它在编译期做的所有优化都建立在“输入形状永远不会变”的假设上。一旦模型在服务端接收的请求长度各不相同静态图要么原地爆炸要么疯狂JIT重新编译编译开销反而成了新的性能瓶颈。这也是为什么很多团队明明上了XLA一跑动态模型还是得开disable_optimizer或者疯狂调allow_growth相关配置。反过来想面对动态shape我们真正需要的其实是一套“按需特化”的机制同一段计算逻辑在shape不同的时候走不同的机器码路径在shape相同的时候能够稳定复用之前生成的特化版本。这套机制和后端表达能力强相关和前端是否静态关系反而不大。把“动态”这个属性放到运行时去解决比在IR层面死磕形状推导要务实得多。我们选字节码VM作为载体不是因为它时髦而是因为它天生适合承载“运行时才知道”的语义。2. 框架设计字节码VM和实时编译怎么结合2.1 为什么是字节码而不是直接LLVM IR说句实话一开始我们也想过直接用LLVM IR作为中间表示。LLVM的优化器成熟、backend齐全、社区活跃看起来是顺理成章的选择。但真正做起来你会发现一个尴尬的事实LLVM IR是给静态编译设计的它对“可变形状”这个概念的抽象能力非常弱。在LLVM IR层面一个tensor就是一个不透明的指针加一段metadata描述shape。你没法在IR里直接表达“这个乘法操作需要循环遍历所有元素并且对shape进行对齐”这种张量语义。你只能靠外部pass提前把shape逻辑全部解析完然后生成一堆纯标量的循环代码。这等于把动态shape的复杂度完全踢给上层编译器而上层编译器一旦碰到循环边界不可静态化的情况就只能生成保守的while循环优化空间基本归零。字节码VM的做法就不一样。字节码天然可以保留高层次的张量语义比如指令集里可以直接有OP_ADD_TENSOR、OP_RESHAPE、OP_GUARD_SHAPE这种指令。解释器看到这些指令可以快速执行JIT编译器看到这些指令可以针对具体的shape做loop fusion、buffer allocation、kernel生成。同一份字节码既能被解释器消费也能被JIT编译器消费语义完全一致不会出现“编译器眼里一套逻辑解释器眼里另一套逻辑”的经典破事。另外还有一个很实际的原因调试体验。直接生成机器码一旦出bug排查成本极高。但如果你把系统分成“前端生成字节码”→“字节码优化pass”→“JIT后端生成机器码”三层每一层都可以单独打印、单测、dump中间状态。我们内部排查shape推断错误的时候直接看字节码的--print-after-all就能定位问题效率比直接看机器码高出一个量级。2.2 整体模块和编译流水线整个系统的骨架分四块前端生成层、字节码优化层、JIT编译层、运行时调度层。前端生成层负责把模型定义不管前端是Python还是C编译成一串字节码。这层不做太重的优化重点是把动态控制流比如循环、条件分支和张量操作之间的依赖关系干净地表达出来。我们选的策略是“一层IR走到底”前端不直接对接LLVM而是生成字节码后续所有优化都在字节码层面完成。字节码优化层跑几组轻量pass死代码消除、常量折叠、shape推断、算子融合。这里的shape推断不是静态编译器那种“我要把所有shape都算出来”而是“能确定的就标注不能确定的就留下符号标记”。比如x 1这种操作无论x是什么shape结果shape都和x一样这种关系可以直接记录在字节码里但x.reshape(-1, 5)这种就只有等到运行时拿到x的真实shape才能决定后续指令。JIT编译层是核心。它接收来自调度层的“编译请求”请求里带着具体的shape签名、dtype、device、layout信息。这一层做的事情有两件一是从字节码生成目标平台的机器码通过内置代码生成器或者降到LLVM二是生成用于运行时分派的guard代码。guard代码的作用是检查下一次输入是否满足已经特化的shape假设如果满足就复用当前编译产物不满足就重新走编译流程。运行时调度层是最容易被低估的一块。它负责所有动态决策维护编译缓存、执行guard检查、调度解释器/JIT两条执行路径、管理workspace内存。调度层设计得好不好直接决定系统在真实负载下的表现。我记得第一次上线的时候JIT编译本身只要30微秒但调度层的哈希计算和锁竞争硬生生把它拖到了300微秒。后来把cache key改成直接基于metadata状态算锁粒度拆细才压回50微秒以内。2.3 一个简单的端到端示例我拿一个非常常见的动态shape场景举例对batch里每条样本做masked softmax。伪代码如下def masked_softmax(x, mask): # x: [batch, seq_len], mask: [batch, seq_len] exp torch.exp(x - torch.max(x, dim1, keepdimTrue)) exp exp * mask return exp / torch.sum(exp, dim1, keepdimTrue)静态编译器看到这函数会抓狂因为seq_len并不是常量。字节码VM的做法是先把它转成一串字节码LOAD x ; 动态shape [b, s] LOAD mask ; 动态shape [b, s] REDUCE_MAX dim1 keepdimtrue ; 输出 [b, 1] SUB ; 广播减法 [b, s] EXP ; 逐元素exp [b, s] MUL mask ; 逐元素乘mask [b, s] REDUCE_SUM dim1 keepdimtrue ; 输出 [b, 1] DIV ; 逐元素除法 [b, s] RETURN这段字节码在执行时调度层会拿到实际的[batch, real_seq_len]签名。如果这个签名之前编译过直接走JIT生成的融合kernel如果没编译过就触发一次实时编译。编译出来的机器码会把上面的多个逐元素操作融合到一个循环里中间不落盘、不产生额外内存。对一个[64, 300]的矩阵融合后的kernel比解释器逐op执行能省下至少3次全量内存读写的开销。这个例子虽然简单但它展示了动态shape下的核心工作方式接口是动态的内核是特化的。3. 核心机制动态shape的特化、guard与缓存3.1 特化粒度怎么选动态shape系统第一个要回答的问题是你到底在什么粒度上做特化粒度太粗比如只特化dtype和device所有shape共享一份通用kernel那JIT基本等于白做性能会和解释器差不多。粒度太细比如把所有shape维度都写死成常量确实能达到静态编译的性能但缓存命中率会惨不忍睹——生产环境里seq_len稍微变一点整个编译缓存就全废了。我们的做法是分层特化信息越稳定特化程度越高信息越易变特化程度越低。第一层特化的是完全稳定的属性dtype、device、维度数量rank、以及每个维度的stride布局是否连续、是否带广播。这些属性在一个模型的生命周期内几乎不会变特化掉它们是纯赚不亏的。第二层特化的是“静态维度”。什么叫静态维度就是在一个计算图中某个维度的值虽然不在编译期可见但从业务角度看它是固定的。比如词表大小vocab_size、隐藏层宽度hidden_size、注意力头数n_head这些维度无论输入怎么变都不会变应该被编译成常量。第三层才轮到真正的动态维度batch大小、序列长度、节点数量这些。这一层我们不做完全特化而是把它们作为“运行时参数”传入kernel。生成的汇编代码里这几个维度作为循环变量存在而不是立即数。这样设计的好处以我实际经验来说缓存命中率可以从70%提升到95%以上。道理很简单一个[64, 300]的输入和一个[64, 310]的输入在细粒度特化的系统里是完全不同的两套代码但在分层特化的系统里它们共享同一套带参数kernel只是运行时传的循环边界不同。热点永远集中在“形状变化但结构不变”的那部分计算上。3.2 cache key设计与guard插桩缓存key是整个调度系统的地基。一开始我们犯过一个典型的错误直接用tensor的指针或者Python对象id作为key的一部分。结果每次torch.empty()新分配出来的tensor即便shape完全一样pointer也不同缓存永远missJIT变成了“每分钟编译一百遍”的悲剧制造机。正确做法是根据tensor的metadata生成签名。签名由以下几部分拼接dtype枚举值device设备类型序号rank维度数量各维度的尺寸动态维度用特殊标记比如-1各维度的stride描述layout连续性必要的特殊属性标记比如是否requires_grad是否sparse我在内部把它称为ShapeSignature。生成签名时要把静态维度直接写成具体数值动态维度写成DYN标记。这样[64, 300]和[64, 310]如果模型把第1维标记为动态签名里就都是[64, DYN]可以命中同一个编译产物。guard插桩也很关键。每次编译一个特化版本时JIT后端会在产物入口处生成一段检查代码验证实际输入是否满足编译时的假设。比如编译时假设了“第二维是动态的第一维必须是64”那guard就检查这两条。如果检查通过跳转到经过特化的快速路径如果检查失败走慢速路径重新调度。guard的语义必须和签名保持一致否则就会出现“缓存命中了但代码跑错”的幽灵bug这类bug非常隐蔽调起来极其费头发。3.3 一个典型编译产物的样子看一段简化的概念性伪代码让没有接触过JIT的读者感受一下特化产物长什么样。假设我们编译上述masked softmax生成代码逻辑大致如下def compiled_masked_softmax(x_ptr, mask_ptr, out_ptr, batch, seq_len): // x: [batch, seq_len], row-major, contiguous // batch和seq_len是运行时传入不是立即数 for b in range(batch): // 反正mask是0/1直接乘在结果上 row_max -inf for s in range(seq_len): row_max max(row_max, x_ptr[b*seq_len s]) sum_exp 0 for s in range(seq_len): e exp(x_ptr[b*seq_len s] - row_max) * mask_ptr[b*seq_len s] x_ptr[b*seq_len s] e // 就地写回省一块临时buffer sum_exp e for s in range(seq_len): out_ptr[b*seq_len s] x_ptr[b*seq_len s] / sum_exp这段代码有三个特征一是所有中间结果都不落临时tensor只用一个标量寄存器流传递二是循环边界全部来自运行时参数不存在“编译期把batch写死”的情况三是对seq_len的访问完全contiguous后端可以安心做SIMD向量化。这种产物既保持了动态性又享受了融合优化正是字节码VM 实时编译的方案应有的样子。4. 字节码设计与JIT实现细节4.1 指令集设计要留什么信息字节码指令集的设计直接决定JIT编译器能拿到多少信息。如果指令集太底层比如一上来就全是load/store/alu那JIT编译器看到的就是一堆碎片化的操作根本没法做全局优化。如果太高层比如一条MATMUL指令包打天下那连基本的融合都无从谈起。我们最终的指令集设计原则是指令面向张量语义元数据面向shape关系。每条指令除了opcode和operand外还会携带一个关系标识告诉运行时这条指令的输出shape和输入shape之间的函数关系。比如指令语义shape关系OP_ADD逐元素加法输出shape 输入shape广播后OP_MUL逐元素乘法输出shape 输入shape广播后OP_REDUCE指定维规约输出shape 输入shape去掉该维OP_RESHAPE形状变换输出shape 编译期推断或运行时确认OP_GUARD运行时断言不产生输出仅用于验证OP_DISPATCH分派到特化版本控制流指令OP_FUSED融合kernel入口由多个op融合而成这个设计解决了两个问题。第一JIT编译器能看懂数据流。它看到OP_ADD和下游的OP_MUL知道“这俩可以融合”因为它们的shape关系一致。第二解释器也能快速执行。解释器不需要理解“融合”这种高阶概念一条一条按语义执行就行。同一份字节码既服务解释器又服务JIT减少了维护两套IR的负担。我需要特别提醒一点指令集里一定要保留shape关系的运算符。比如broadcast和align不能只体现在shape元数据层面还要有对应的操作符或明确的shape传播规则。否则JIT编译器在生成融合循环时就得自己去猜广播逻辑猜错一次就是shape不匹配的bug非常难排查。4.2 融合与内存分配的细节融合是JIT编译性能提升的最大来源之一。对于elementwise运算链比如exp、mul、div这类逐元素操作我们会做水平融合把多个独立的op合并成一个并发循环。这样做的收益是可以直接量化的一个[1024, 1024]浮点矩阵一次全量读写大约是16MB内存流量如果一条链上有4个op解释器要产生4次读4次写而融合kernel只要1次读1次写访存流量直接砍掉75%。做融合的时候有个细节很容易踩坑有view操作的节点不能随便融合。比如x.reshape(y)和y * 2如果你试图把reshape折叠到前面的乘法循环里就要保证reshape不会产生实际数据搬运。好在我们指令集里明确区分了OP_VIEW纯元数据操作不搬数据和OP_COPY物理拷贝JIT编译器看到OP_VIEW链就直接忽略看到OP_COPY就作为融合边界。内存分配这块是动态shape系统最容易变成性能瓶颈的地方。编译时不知道shape就不能预先精确分配中间buffer常见的做法是运行时动态计算所需大小然后用内存池分配。我们直接做了个arena内存池按编译产物的buffer需求列表一次性从池里取一块区域kernel执行完整体归还。这样避免了每个op单独malloc/free的碎片化开销区别非常明显——同一个masked softmax用arena比每次现分配临时tensor整体耗时能缩短20%以上。4.3 编译延迟怎么压实时编译最怕的是编译时间吃掉计算时间。如果编译一个融合kernel要花50毫秒而计算只要0.5毫秒那整个方案就是一场灾难。所以编译延迟是必须死磕的指标。我们做了三件事来压低编译延迟第一分级别编译。不是所有输入都走完整的机器码生成路径。对于特别小、特别频繁的shape比如固定batch、固定seq_len的场景我们直接把字节码解释执行的开销压到最低只有确认“这个shape会反复出现”才升级到编译路径。判断依据是缓存里的命中频率计数器命中超过一定阈值就触发编译。这个设计叫lazy promotion很好地平衡了冷启动和热稳定。第二代码生成设快速通道。机器码生成不直接走通用LLVM优化流水线而是先走一个轻量的快速通道只做循环融合、循环展开、常量替换不做昂贵的全局优化。快速通道生成的代码质量虽然比不了全优化版本但胜在编译速度快一个中型kernel 20~30微秒就能出结果。等这个kernel被反复命中后后台再异步提交一次全优化编译下一次命中时换成更优的版本。这个“先能跑再跑好”的节奏非常实用。第三复用编译缓存池。编译产物不但缓存在内存里还按device分别隔离。同一个shape签名在CPU和GPU上分别生成不同代码互不干扰。同时缓存池用LRU策略淘汰避免某些大batch的场景占满全部缓存。编译延迟这件事我的经验是一微秒一微秒地抠是值得的。因为生产环境里动态shape是常态每一次shape变化都可能触发一次编译编译延迟直接叠加在用户的请求延迟上。快速通道定在30微秒这个量级对绝大多数推理场景都是可接受的。5. 踩坑记录常见问题与排查心得5.1 我撞过的六个典型问题第一个问题是缓存永远miss。一开始我们把cache key设计成包含tensor对象指针结果每次新的分配都是新的指针缓存命中率接近于零。这个问题排查的时候也比较恶心因为JIT编译本身没有报错只是执行时间忽高忽低像是随机抖动。定位到以后把key改成了基于metadata的签名命中率立刻恢复正常。排查经验看到“性能忽好忽坏”优先怀疑cache key。第二个问题是guard过强导致频繁重编译。我们早期对layout的guard要求所有维度必须contiguous结果碰到转置操作视图就直接触发重新编译。后来放宽了guard条件允许输入是带stride的view但要求代码生成时把stride信息传进去这样非连续的view也可以安全执行。这是个典型的“以性能换缓存命中率”的trade-off实际收益远大于损失。第三个问题是多线程并发编译导致重复工作。服务工作线程同时达到两个相同shape的请求两个线程同时发现cache miss同时开始编译同一个kernel白费一倍的编译成本。解决方法是给cache加per-key的编译锁第二个线程发现key不存在时会等待而非重新编译等锁释放后直接复用第一个线程的结果。这个改动很小但编译期间的系统CPU占用率立刻下来一大截。第四个问题是动态循环边界导致vectorization效果差。我们把循环边界作为运行时参数后LLVM后端判断不出循环上界是否对齐SIMD向量化自动降级成标量循环。后来在快速通道里手动做循环剥离主循环写成对齐的SIMD循环尾部处理标量尾数。这样做以后浮点峰值利用率提了大概40%。第五个问题是融合后数值结果不一致。做exp/mul/div融合时因为寄存器里直接传中间值不再落到内存浮点计算的中间舍入路径发生变化导致结果和逐op执行版本在最后几位上存在差异。这个问题不涉及bug但从用户视角看就是“结果变了”。我们的做法是配置一个strict_ieee开关默认开启确保数值一致性只有在明确要求性能并且用户接受微小误差时才关闭。第六个问题是workspace内存碎片化。动态shape下buffer大小变化频繁反复new/delete导致内存碎片一个原本几十MB的workspace实际占用能做到两三百MB。这块没有特别优雅的解法最终就是用arena pool按大小分级复用效果很明显峰值内存占用降了约一半。5.2 问题速查表我把上面这些踩坑经验整理成一个速查表方便团队新同事快速对齐现象可能根因定位手段解决方案性能忽高忽低cache key设计不当打印cache hit/miss日志改为metadata签名频繁重编译guard条件过强dump guard失败原因放宽guard传strideCPU编译负载高并发重复编译检查编译请求去重per-key编译锁SIMD利用率低动态循环边界查看生成代码内层循环循环剥离尾部处理结果略有差异融合改变舍入路径对比逐op与融合输出strict_ieee开关注释内存膨胀workspace碎片化统计内存分配请求arena pool按大小复用这张表里的每一条都是从实际故障里长出来的。我个人的体会是系统问题很少是单点的往往是cache、guard、内存、并发几个维度互相纠缠所以排查时一定要先分层别一上来就盯代码生成器。6. 实测效果与个人经验6.1 我在真实任务上的测量我们在两个代表性场景上做了评测一个是变长序列的Transformer推理一个是GNN邻居采样训练。前者是典型的batch内长度不一致后者是典型的运行时才确定的动态shape。变长Transformer场景里基线是用padding把所有batch补齐到最长序列按静态shape编译。我们这套字节码VM方案以动态shape直接编译运行不padding。同样的batch融合后的kernel实际只处理真实长度访存量下降约35%端到端延迟提升在15%~25%之间具体数值取决于batch内长度的方差——长度越参差收益越明显。GNN场景更夸张一些因为每个节点的邻居数不规则静态编译根本没法落地。我们只能跟解释器基线对比启用JIT编译后aggregation相关的kernel执行时间缩短了约4倍主要来自两个op的融合和循环参数的动态化。编译本身的额外开销单次编译平均在40微秒左右对于一个batch耗时动辄几毫秒的训练步来说几乎可以忽略。我还特意测过一个持续变化的极限场景每个请求的seq_len在32到512之间随机变化。这种情况下缓存命中率是衡量系统健康度的核心指标。分层特化设计跑出来的命中率稳定在94%以上意味着每100个请求只有不到6个需要重新编译。这个数据让我对这套设计有了信心动态是常态但底层计算模式几乎没有变化只要把“变化的维度”参数化缓存就能持续发力。6.2 做这类系统最重要的几条工程心得第一设计方案之前先想清楚你的“动态”是哪个维度上的。是维度大小动态还是结构动态比如有没有哪条分支可能不存在还是控制流动态比如循环次数可变这三种动态对系统的影响完全不同。把这个问题想清楚后面所有设计都会顺利很多想不清楚就容易做出一个“既要又要”的怪物。第二编译器和运行时调度必须作为一个整体设计。我见过不少项目编译器和调度器分开两个团队做结果编译产物暴露给调度器的接口太粗糙调度器没法知道“这个kernel适合什么样的输入”于是要么保守地总是走解释器要么激进地总是重编译。这个接口的设计比后端优化重要得多。第三别追求极端特化追求汇合点。极端特化每个shape生成一份代码性能顶天但缓存崩盘。完全不特化缓存总能命中但性能平平。工程上要追求的是“大部分时间在特化路径上运行小部分时间在通用路径上兜底”的状态。我实际体会是优化方向应该盯着“平均延迟”和“p99延迟”而不是盯着“理论峰值”。第四调试工具在第一天就要建。能够dump字节码、dump生成的IR、dump cache状态、统计compilation次数和命中率这些能力在系统早期不值钱但后期遇到莫名其妙的问题时每一分钟调试时间都能被这些工具回本。我甚至建议直接把编译日志设计成可以回放的形式这样线上问题可以离线复现。最后分享一个经验性判断动态shape计算不是某个框架特有的事而是所有生产级推理系统都绕不过去的坎。与其在静态编译器里不断打补丁不如从运行时的角度重新思考——字节码VM加实时编译算是我验证过的、比较务实的答案。如果你也在做类似方向我的建议是先从最痛的一个场景切入把端到端链路跑通再逐步扩展指令集和融合规则不要一上来就想着做一个万能编译器。

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

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

免费获取报价 →
↑