
1. 动态张量为什么天生和编译优化不对付1.1 一个真实的性能现场变长序列把GPU拖垮了大概半年前我在优化一个变长序列的推理服务。那批数据每条样本长度差异非常大短的只有十几个token长的能到几百。为了跑batch常规做法是pad到最长长度于是你会在性能面板上看到一组非常刺眼的数据显存翻了三倍但GPU利用率只有30%出头CPU倒是快跑满了。用火焰图往下钻开销根本不在计算上而全耗在张量形状检查、中间缓冲区分配和Python到C的反复跨语言调度上。这个现象在动态张量计算场景里非常典型。所谓动态主要指三个维度第一是形状动态输入张量的batch、序列长度、特征维度在运行时才会确定第二是控制流动态比如while循环、条件分支要等实际数据出现才能走通第三是内存布局动态同一个逻辑张量会因为切片、转置、拼接不断改变stride和存储位置。静态编译器拿这种输入没太多办法纯解释器倒是能应付但性能又差得不像话。于是我开始系统性研究一个折中方案先做一个字节码虚拟机把张量程序编译成字节码跑起来之后再对热点路径做实时编译。这套东西的核心思路和LuaJIT有点像但优化对象从标量运算换成了张量计算要解决的问题也从解释器太慢变成了动态形状导致的重编译风暴和kernel启动开销。1.2 静态编译在动态形状面前的挣扎静态编译器和动态张量之间有一条很深的鸿沟。以T个时间步的循环为例静态编译器在第一次看到输入shape是(batch4, T7, 512)时可以很舒服地做特化编译把T7直接折叠进循环上界allocate固定大小的中间缓冲把该融合的算子全部融合掉。但第二次请求T变成了9整个特化全部作废又要重新编译一遍。如果生产环境里T的取值分布有二十多种结果就是二十几份特化代码互相换血每次切换都伴随一次几百毫秒甚至秒级的编译停顿。更难受的是编译器在编译期如果发现某个维度的值不可知通常只能退化成非常保守的代码生成很多融合优化当场失效性能甚至不如写得不怎么好的eager代码。统计下来这类系统在动态输入下的真实编译命中率往往不到50%——也就是说一大半时间在做无效编译。我在一个用TVM做的实验里验证过当序列长度集合是{3,4,5,6,7,8,9,10}时XLA风格的静态编译几乎每个长度都要单独处理中间切换的开销大于收益最终端到端只比eager快了一点点完全对不起编译带来的复杂度。1.3 解释器慢而灵活编译器快而僵硬中间路线怎么走纯解释执行和纯静态编译是两个极端。解释器的好处是天然支持动态控制流、动态形状遇到什么shape就处理什么shape缺点是每条张量指令都要走一遍大循环做完结果还得写回内存运算密集度一高就完全发挥不出硬件性能。我手写的块解释器跑一个简单的矩阵乘法链比起直接用cuBLAS/oneDNN差了可能有四五倍——大量时间耗在指令分派和中间张量的读写上。实时编译JIT编译正好卡在中间程序先用字节码形式解释执行解释器同时记录每一条指令经过时的形状分布、调用频率某个字节码位置被发现是热点且形状模式稳定就把它编译成高度特化的机器码。以后再来同样的形状直接命中缓存发现形状变了要么编译一个新版本要么先退回解释器兜底。这套架构的关键词就三个字节码VM提供灵活的执行基础profiler识别热点JIT编译器把热点路径变成特化代码。训练和推理都能覆盖推理场景因为形状分布往往更集中收益尤其明显。2. 指令集与执行模型先把张量字节码设计对2.1 两层指令集控制流与张量运算分离设计字节码VM的第一步是定义指令集。我踩过的坑是别把张量指令和控制流指令混在一起设计。张量指令的语义很重每个操作都牵扯到shape推断、dtype检查、设备分配而控制流指令是轻量级的它们在解释器里几乎不产生实际计算。混在一起会让dispatch逻辑变得极慢。我最终分成了两层层次指令示例职责Flow层jump、branch、call、ret、loop_head管理PC跳转、调用栈、循环状态Tensor层matmul、add、reshape、transpose、layer_norm真正的张量运算带完整shape语义一个动态循环累加的程序编译出的字节码大概是这样的# 伪代码: sum 0; for i in range(T): sum x[i] load.const %0, 0 # sum 0 load.tensor %1, x # 加载输入张量 load.const %2, 0 # i 0 loop_head: # 动态上界判断 get.shape_axis %3, %1, 0 # 得到第0维长度 T cmp.lt %4, %2, %3 # i T ? branch.false %4, exit # 取 x[i]这里x如果是连续存储可以做指针步进 tensor.index %5, %1, %2 tensor.add %6, %0, %5 move %0, %6 add.const %2, %2, 1 jump loop_head exit: ret %0注意loop_head里的get.shape_axis这是在动态形状下保持灵活的根基。每次进入循环体都重新读一次当前轴的长度这样即使运行期间张量被resize过循环也不会跑错。代价是每轮多一条shape读取指令但相比重新编译整个循环这点开销完全可以接受。2.2 张量值表示shape信息要能参与运行时计算张量的值表示是VM设计里最容易想简单、又最致命的一环。最朴素的表示是void* data shape stride dtype但这在动态场景下不够用——你还需要一张shape的自由度图告诉VM哪些维度是固定的、哪些维度是符号化的符号与符号之间有没有约束关系。我最后采用了带符号shape的TensorValue结构struct TensorValue { void* data; std::vectorint64_t shape; std::vectorint64_t stride; DType dtype; Device device; // 符号shape-1表示未知0表示已知 std::vectorint64_t sym_shape; // 符号约束表比如 sym[0] sym[1] };sym_shape数组是关键。假设两个矩阵A、B要相乘A的shape是(M, K)B是(K, N)如果M、K、N都是-1未知VM不会立刻报错而是记录一条约束关系A.shape[1] B.shape[0]。等运行到matmul指令时再拿实际shape去匹配约束。匹配成功就正常计算匹配失败才算真正的运行时类型错误。这样设计的好处是解释执行和JIT编译可以共享同一套shape推断逻辑。2.3 为什么选寄存器式而不是栈式如果你做过脚本语言VM一定纠结过寄存器式和栈式。栈式的优点是字节码紧凑、编译简单Java VM、Python VM都走这条路。但张量运算的特点是每个操作都涉及大块显存/内存的读写如果每条指令都从操作数栈顶弹出再压入会频繁产生中间TensorValue的引用计数切换最终表现为大量的引用计数原子操作和shallow copy。寄存器式VM天然接近于SSA形式每条指令明确写出dst, src1, src2解释器可以直接把src的指针传给底层kerneldst的输出位置处在一个可预测的寄存器槽里。后续要做JIT优化时这种形式几乎可以直接映射到LLVM IR或者等价的三地址代码不需要再做一遍phi节点重建。代价是字节码体积变大、编译器需要做寄存器分配。对于张量程序来说一块显存缓冲比一条字节码指令值钱得多这个tradeoff非常划算。实际上TVM的Relax早期设计也走了类似路线核心IR是函数式带显式shape标注的只是它没有硬性地落到一个独立的VM指令集上。3. 实时编译三件套形状特化、指令融合与后端代码生成3.1 形状特化与多态版本缓存实时编译最核心的是形状特化。流程是这样的解释器在运行每个字节码基本块时维护一张频率表记录该基本块被执行了多少次、进入时张量形状的签名比如M32,K512,N128。当频率超过阈值——我常用8次——就把这个基本块标记为可编译块并尝试针对那个形状签名编译一份特化版本。编译出的特化版本放入一个形状签名到编译后代码的哈希表。后续执行到这个字节码位置时先查一下当前的实际形状签名在不在哈希表里。命中直接跳转到特化版本没命中就留在解释器继续跑同时启动一次异步编译。这里有一个很值得说的细节特化版本并不只对应一个形状。比如MatMul的形状签名是(M,K) x (K,N)K在矩阵乘法里是归约维度K一样但N变化很多kernel可以复用底层BLAS选择结果参数化一下就行。可以把签名分成强特化组和弱特化组强特化维度变了必须新编译弱特化维度变了只需修改kernel参数。这个分组能把编译缓存命中率提升得非常明显我在实验里中位数shape命中率从53%拉到了87%。3.2 指令融合把一串字节码变成内核有了特化代码接下来就是融合优化。举一个layer_norm的例子字节码序列一般是tensor.mean %1, %0, axis1 tensor.sub %2, %0, %1 tensor.square %3, %2 tensor.mean %4, %3, axis1 tensor.add.const %5, %4, eps tensor.sqrt %6, %5 tensor.div %7, %2, %6如果每一条都单独启动一个CUDA kernel显存里会多出五六个中间缓冲区而且每个kernel启动都有几十微秒的固定开销。实时编译器要做的就是把这条链检测出来融合成一个巨kernel把mean的结果算在寄存器/smem里不落显存把square直接折叠进sum最后rescale一次。融合边界要非常小心。控制流指令绝不能跨过去因为融合后的kernel必须是直线代码依赖外部位点比如某些全局参数的指令也必须分开否则缓存后的特化代码在参数变化时得不到正确结果。我的经验是编译器内部先构建一张数据依赖图只做纯数据流子图的融合一旦遇到side-effect指令就立刻截断。3.3 代码生成后端LLVM、TinyCodeGen还是自研实时编译的后端选择直接影响工程复杂度和编译延迟。我实验过三条路线LLVM OrcJIT优化效果好能吃到AVX-512、SVE这类新指令集的红利。但编译延迟高冷启动一个模块经常要几十毫秒到上百毫秒对于几十毫秒一次的前向推理来说等不起。生成C代码再调系统编译器用模板生成一段C代码丢给cc -O2 -shared编译成.so再dlopen。有效但每一步都要起子进程延迟同样难以控制。轻量级自研代码生成器只针对自己指令集里的几十条张量指令直接输出内存中的机器码。编译速度快几微秒到几百微秒搞定缺点是优化能力有限得靠更高层的融合逻辑兜底。最终线上系统选的是第三条配合一层基于指令模式匹配的融合优化器。效果是编译延迟从30ms压到了2ms以内融合后的内核性能比未融合的特化版本提升平均约2.3倍。如果你的系统里没有特别激进的新指令集需求这个组合是性价比最高的。4. 字节码VM在动态张量战场的差异化定位4.1 和PyTorch系方案比差在哪、好在哪里PyTorch的TorchScript和TorchDynamo也存在解决动态形状问题但路径本质上仍然是图模式先抠出一个静态计算图再交给后端编译。TorchScript的做法是尽量追踪出一个静态图遇到动态控制流就抽象成loop/branch子图TorchDynamo则靠guard机制每次进入Python函数时检查shape是否和上次一致不一致就重新编译。TorchDynamo的问题在于guard机制是全函数级的。函数里只要有一个维度是动态的guard就会频繁失效导致cache miss和重编译风暴。我在一个多层GRU的动态batch训练实验里Dynamo的cache miss率一度达到40%——每次miss都要回落到eager模式重新trace性能反而不如干脆全程eager。字节码VM的优势在于特化粒度是基本块级而不是函数级只有真正热点、形状稳定的基本块才会被编译其他部分留在解释器里互不干扰。4.2 XLA和TVM面对动态形状的挣扎XLA在动态shape上的核心手段是Pad和DynamicReshape思路是把动态shape的输入pad到某个上限让kernel逻辑变成静态的。这个方案在推理场景还勉强能用但训练场景会遇到问题pad出来的部分是无效计算反向传播时还得专门mask掉计算浪费可达30%-70%。而且一旦用户给的上限不准确超出部分会直接报错非常不友好。TVM的Relax尝试在IR层面引入shape表达式把形状本身变成一等公民参与编译决策。思路很优美但工程复杂度过高。我当时评估过自己搞一套Relax-like IR要多久——排期下来一个四人小组至少四个月。字节码VM这条路就务实得多不追求一次成功而是解释执行保证正确缓存命中保证性能上线负担小很多。4.3 渐进式执行不编译也能活编译只是加速字节码VM最独特的地方在于它的渐进式特质。这意味着系统上线不存在要么全有要么全无的风险。如果生产环境的形状分布极端发散特化缓存命中率上不去系统会自动退化成纯解释执行性能顶多比eager慢个20%绝不会出现编译风暴把服务打挂的情况。也正是这种特性让我敢把这套VM直接用在线上推理服务里。它像一个带涡轮增压的引擎涡轮不介入的时候就是普通自然吸气一旦转速到位涡轮介入就能获得明显性能提升。对于动态张量计算这种大部分时间形状稳定、偶尔有长尾的负载没有比这更合适的形态了。5. 上线前必须啃下的工程硬骨头5.1 shape推断在动态场景里的隐藏大坑动态shape下做编译优化最脏的活永远是shape推断。难点在于当两个符号维度看上去可能相等、但编译器无法在编译期证明时你只有两条路放弃优化或者加入运行时断言。我踩过一个特别典型的坑计算Attention Score时Q和K的sequence长度在绝大多数情况下相等因此在编译期直接假设了它们相等做了内存布局的合并。结果某天线上来了一条极端数据Q长度和K长度不匹配程序直接segfault。修法是给这个假设加一条运行时shape断言assert_eq(q_len, k_len)失败就跳转到通用版本重算而不是继续跑特化代码。从那以后我给自己定了一条规矩动态shape的特化编译必须携带约束表约束表验证失败就回退宁可不优化也不出错误结果。5.2 副作用指令会让特化代码状态错乱张量计算里有不少带副作用的操作in-place update、随机数生成、显存池的显式分配释放、某些算子内部的全局状态。这些指令如果被无脑融合进特化版本会导致编译后代码和解释器执行结果产生分歧。举一个真例。某个模型在循环体里反复调用随机dropout我最初把dropout和后面的matmul融合成一个kernel看起来没问题。但训练时梯度下降需要dropout的mask是可复现的——不同batch的同一位置要产生确定性mask。融合后由于运算顺序变化串行随机流被打乱复现实验直接失败。最终的做法是所有带随机状态的算子都不参与融合保留独立的kernel调用只在它们周围做shape特化。5.3 编译失败后的回退要静默且稳妥再健壮的编译器也有veto的情况某个shape组合超出发射器的处理范围、某些指令还没移植到JIT后端、或者运行时资源不足。我见过很多项目在这些场景下直接抛异常导致整条推理链路崩溃。我的做法是给每个编译单元一个编译失败计数连续失败三次就永久标记为不编译让这段字节码永远跑解释器并输出一条warning级别的日志。回退策略的另一个要点是回退必须在字节码级别的同一位置进行不能出现前半段跑特化、后半段跑解释器这种割裂状态。所以我在JIT入口设计了一个两阶段提交先验证当前状态与编译时的假设完全一致再切换PC到特化代码的入口任何一步验证不过都留在解释模式。5.4 编译器线程不能拖慢推理主线程实时编译的延迟会直接影响端到端延迟尤其是推理场景。我的调优心得是把编译完全放到独立线程池里主解释器只做发布编译任务和查询编译结果。解释器每到达一个热点块先向编译线程池提交任务然后继续解释执行编译完成后结果被放回缓存区下次执行到该位置时再启用。还需要注意编译器线程池大小和CPU核心数的关系。编译过程是CPU密集的如果线程池设太大会和主线程抢核导致推理延迟反而上涨。我的经验值是主服务部署在16核机器上线程池设4既能并行编译多个shape版本又不至于把CPU吃满。实践中还可以结合profile阶段的采样频率来动态调整线程池大小。6. 实测数据与调优心得6.1 变长序列Transformer推理从45ms降到28ms拿一个经典的变长序列Transformer decoder推理做基准输入batch32样本长度在5到60之间变化使用RTX 3090。同一模型分别在eager模式、纯解释器VM、字节码VM实时编译三种状态下跑各跑5000次取中位数延迟执行模式中位延迟相对EagerEagerPyTorch原生45ms1.0x纯解释器字节码VM58ms0.78x慢字节码VM 实时编译28ms1.6x快解释器慢是预期内的它存在的意义本来就是兜底和采集profile。实时编译带来的提升主要来自两块约65%的算子被融合成巨kernelkernel启动次数从每层十几次降到了两三次另外形状特化后能提前锁定内存池里具体大小的缓冲块避免动态分配这部分省了约20%的时延。6.2 动态batch的RNN训练编译预热期的等待训练场景的时间线值得单独记录。前8个step基本都在解释执行profiler在后台收集基本块的形状分布。第9个step开始第一个隐藏层基本块的特化编译完成之后每隔几个step就多一份特化版本。从第25个step开始所有热点基本块都编译完成单step时间从解释期的18ms降到了8ms加速比约2.2倍。需要提前做好心理准备的是编译期的前30个step会有明显波动某些step甚至会因为后台编译器线程抢占CPU而比解释期还慢。解决方法是把编译线程优先级调低同时把是否开启实时编译做成运行时开关在训练的前几个epoch先纯解释执行收集统计等形状分布稳定了再打开开关。6.3 几个直接影响收益的参数最后分享几个调参经验都是我用控制变量法试出来的热点阈值我默认8次。阈值太高小batch请求还没触发编译就结束了收益为零阈值太低只有两三次调用也去编译编译开销比解释执行还大。在推理场景可以降到3-4次因为热点的形状分布通常更集中。特化缓存容量默认1024个shape版本LRU淘汰。容量太小长尾形状频繁挤掉核心版本容量太大哈希查找本身开始有可察觉开销。1024够覆盖90%以上的线上shape分布。禁止编译名单对于延迟极度敏感、且形状分布确实非常发散的场景直接禁用实时编译走纯解释反而更稳定。我写了一个运行时自适应开关连续20次查询到特化缓存miss且编译完成率低于30%就自动降级为纯解释模式。我个人在实际操作中的体会是这类系统最忌讳一上来就追求完美编译应该先把字节码解释器跑通让业务正常运作再把profiler和JIT逐步加上去。每加一层优化都盯住端到端延迟的P99和缓存命中率一旦发现收益为负就果断撤掉用数据说话。动态张量的优化没有银弹但字节码虚拟机加实时编译的组合确实是目前我在生产环境里验证过、最能在灵活性和性能之间取得平衡的方案。