☰
动态shape下的张量编译:字节码VM如何用实时JIT突破性能瓶颈
2026/10/5 1:03:27 网站建设 项目流程

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, dim=1, keepdim=True)) exp = exp * mask return exp / torch.sum(exp, dim=1, keepdim=True)

静态编译器看到这函数会抓狂,因为seq_len并不是常量。字节码VM的做法是先把它转成一串字节码:

LOAD x ; 动态shape [b, s] LOAD mask ; 动态shape [b, s] REDUCE_MAX dim=1 keepdim=true ; 输出 [b, 1] SUB ; 广播减法 [b, s] EXP ; 逐元素exp [b, s] MUL mask ; 逐元素乘mask [b, s] REDUCE_SUM dim=1 keepdim=true ; 输出 [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也不同,缓存永远miss,JIT变成了“每分钟编译一百遍”的悲剧制造机。

正确做法是根据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,传stride
CPU编译负载高并发重复编译检查编译请求去重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加实时编译,算是我验证过的、比较务实的答案。如果你也在做类似方向,我的建议是先从最痛的一个场景切入,把端到端链路跑通,再逐步扩展指令集和融合规则,不要一上来就想着做一个万能编译器。

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

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

立即咨询