
做推理引擎这些年我最大的感受是静态 shape 的优化已经卷到头了真正拉开差距的反而是动态 shape 的支持能力。前阵子我们服务里一个模型输入尺寸从固定 512 改成允许 256 到 1024 动态变化原本编译好的计算图直接报废离线编译的二十多个 shape 变体全都不匹配线上流量一上来就卡在 shape check 上。后来我们把精力砸在“计算平台元数据定义”这一层重新设计了 shape 表达式的表示和编译期的约束传播才算把动态 shape 的坑填平。这篇文章就聊聊我在这过程中的理解计算平台里的元数据到底在定义什么计算图编译为什么怕动态 Shape以及一套能落地运行的动态 Shape 支持机制应该怎么设计。内容偏底层适合正在做推理引擎、编译器优化或者被动态 shape 折磨过的小伙伴看完可以直接对照自己项目里的数据结构和技术选型。1. 先从元数据定义说起1.1 一个 Tensor 的身份证到底包含什么很多人一听到“元数据定义”第一反应就是 shape 和 dtype。但在真正的计算图编译流程里Tensor 的元数据远不止这两个字段。我常用“身份证”来打比方shape 和 dtype 只是姓名和性别真正决定你能不能坐高铁、住酒店的还有更多字段。一份完整的 Tensor 元数据在编译器和运行时之间传递时至少要包含以下几类信息字段作用影响范围shape每个维度的大小静态时为常量动态时为符号表达式内存分配、kernel 选择、算子融合dtype元素类型如 fp32、fp16、int8算子实现、精度策略、内存对齐strides每个维度相邻元素间的步长是否连续、是否需要 layout 转换memory formatNCHW、NHWC、NC1HW2 等卷积优化、向量化访存layout infotensor 在寄存器/共享内存/全局内存的分布计算与访存重叠quant/scale量化参数、zero point量化算子融合、反量化位置生命周期临时变量还是持久变量是否可复用内存规划别小看 strides 和 memory format 这两个字段。一个x[:, ::2]的切片结果shape 看着没问题但 strides 是乱的如果编译器没捕获这个信息后续直接把该 tensor 当连续内存传给 kernel数据全错。动态 Shape 问题也是一样它不只是把 shape 从常量变成变量而是会让上面所有字段的推导从“编译期可确定”变成“运行期才可确定”。所以在设计元数据的时候我强烈建议把 shape、dtype、strides、memory format 看成一张完整的表而不是单独处理。动态 Shape 支持机制本质上要保证这张表在“未知维度”存在的情况下依然可以被推导、被约束、被验证。1.2 为什么静态 Shape 这么好用静态 Shape 指计算图中每个 Tensor 的每一个维度在编译期都是确定的整数。这种情况下编译器的日子非常好过因为所有优化 pass 拿到的都是精确数字可以做很多激进的重构。举个例子给定输入[1, 3, 224, 224]编译器可以直接算出卷积输出是[1, 64, 112, 112]然后提前分配一块固定大小的内存所有中间结果都在这块 arena 里复用。kernel 选择也简单224 的尺寸配 3x3 卷积用滑窗实现还是用 implicit GEMM按 shape 大小查表就能定。还有 layout assignmentNCHW 还是 NHWC编译器可以根据后续算子的访问偏好统一决定。这种模式在目标检测、分类这种固定分辨率的模型上非常有效很多部署栈也因此把“静态 shape”当成默认配置。但现实场景里动态 shape 是躲不掉的NLP 模型里 batch size 和序列长度几乎不可能固定一条样本 32 个 token另一条 512 个 token模型结构一样shape 完全不同。推荐系统里用户特征数量、候选 item 数量都是动态的。目标检测模型经常要接收不同分辨率的输入或者在一个 batch 内做 multi-scale 推理。静态 shape 编译那一套在动态输入面前很尴尬你没法提前知道内存多大没法固定 kernel 参数连计算图里的中间 tensor 是否存在都会变。于是就有了“动态 Shape 支持机制”的需求。2. 动态 Shape 为什么让编译器头疼2.1 动态 Shape 和动态图不是一回事先把概念理清很多人会把“动态 Shape”和“动态图”混在一起。动态图是 PyTorch 默认那种 eager 模式每次执行一行 Python 代码就现场建一个算子节点现场执行。这种情况下 shape 是天然动态的因为根本没有一张先编译后执行的完整图也没有所谓的“编译期优化”。而计算图编译模式下我们希望先把整张模型编译成一个可执行文件或可部署的二进制包再在服务端反复执行。这时候输入 shape 如果是动态的编译器面对的问题就是如何在不知道具体维度数值的情况下生成正确且高效的代码。所以动态 Shape 支持机制是静态图和动态语义之间的一个妥协静态编译的底子还在但要给“未知量”留出表达空间。这个“未知量”就是符号 dimension所有的难点都围绕它展开。2.2 一个动态维度引发的连锁反应表面上看动态 shape 不就是一个维度写“None”嘛运行时拿到具体值填进去不就行了但真正做一次 shape 推断就会发现一个动态维会像多米诺骨牌一样把下游所有 Tensor 的元数据全部变成表达式。假设输入是[N, L]其中 N 是 batchL 是序列长度二者都是动态的。经过 embedding 后变成[N, L, 128]再经过reshape变成[N * L, 128]然后做矩阵乘。下游输出 shape 写成[N * L, 64]还算简单可一旦中间夹着切片、concat、split 或者带 mask 的操作shape 表达式就会越来越复杂。比如# 伪代码 y x[:, :seq_len] # [N, seq_len, 128] z concat([y, padding], axis1) # [N, max_len, 128] w z.reshape([N * max_len, 128]) # [N*max_len, 128]这里的seq_len是一个动态值max_len又可能是另一个约束下的常量。编译器要回答的问题变成我怎么知道z的第二个维度等于seq_len padding_len我怎么证明 reshape 操作的元素数量是可整除的我在内存上应该给z预留多少空间这些问题如果全靠运行时去碰运气那编译优化基本无从谈起。所以动态 shape 难的不是“报错”而是如何在编译期留下足够多的元数据约束让编译器即使不知道具体值也能做出安全、不退回解释执行的决定。2.3 静态优化手段集体失效一旦 shape 变成动态以下几类优化会遇到大麻烦布局优化NCHW 和 NHWC 的转换往往依赖固定维度顺序。动态 shape 下维度的顺序本身是固定的但每个维度的“语义”在运行期才绑定到具体值布局 pass 不能随手改内存 order否则运行时校验会很麻烦。内存复用arena 方式要求所有 tensor 的 size 编译期可算。动态 shape 下每个 tensor 的 size 是一个表达式两个 tensor 能不能复用同一块内存需要做表达式等价性判定这比整数比较难得多。算子融合某些融合要求中间 tensor 的 shape 是静态的比如卷积ReLU 可以融合但卷积动态 reshape 就不能乱来因为 reshape 可能在处理动态维时引入数据拷贝。kernel 选择GEMM 有专门为小 shape 和大 shape 优化过的实现编译期没法确定走哪条分支运行时再做 dispatch 又可能带来额外开销。动态 Shape 支持机制本质上就是为这些“失效点”逐一提供替代方案内存上做动态规划kernel 上做多版本派发融合上尽量收缩动态区域。3. 元数据定义层怎么设计才能扛住动态 Shape3.1 从 int list 升级到 shape expression既然动态 shape 无法用固定整数表示第一步就是改造 Tensor 元数据里的 shape 字段。我见过最偷懒的做法是把某个维设成0或者-1来表示动态然后在运行时去填。这个方案能做 demo但完全无法支撑编译优化因为-1不携带任何约束编译器对它的了解为零。更好的做法是把 shape 设计成“表达式”而不是一个简单整数。每个维度可以是下面三种之一常量整数例如128符号变量例如N、seq_len符号变量的表达式例如N * 2 4在代码里大致长这样from dataclasses import dataclass from typing import Optional, Union, Tuple dataclass(frozenTrue) class SymDim: name: str lower_bound: Optional[int] None upper_bound: Optional[int] None dataclass(frozenTrue) class TensorMetadata: shape: Tuple[Union[int, SymDim, tuple], ...] # 简化版tuple代表表达式 dtype: str strides: Optional[Tuple[Union[int, str], ...]] None memory_format: str dense quant_scale: Optional[float] None quant_zero_point: Optional[int] None这里我解释一下为什么用符号变量而不是直接写死表达式节点。符号变量的意义在于它可以在整个计算图里共享。同一个 batch size 出现在输入、中间结果、输出的多个维度里它们应该引用同一个符号这样编译器才能知道x1.shape[0]和x2.shape[0]是同一个值可以做维度绑定。表达式节点则可以表达缩放和偏移比如某种 padding 后的序列长度是seq_len 2这比新造一个符号更准确也更利于约束传播。3.2 在 IR 级别表达动态 Shape光在宿主语言里定义 TensorMetadata 还不够计算图编译的各个 pass 需要在中间表示IR里直接看到动态 shape。否则每个 pass 都要绕回到外部数据结构去查工程上很容易失控。我常用的做法是在 IR 的 type 体系里直接支持“符号 shape”。如果你们用的是 MLIR 风格可以写成类似这样func.func inference(%arg0: tensor?x?xf32 {shape_sym [N, L]}) - tensor?xf32 { %c128 arith.constant 128 : i64 %slice tensor.extract_slice %arg0[0, 0] [%N, %L] [1, 1] : tensor?x?xf32 to tensor?x?xf32 %pad tensor.pad %slice low[0, 0] high[%c128, %c128] : tensor?x?xf32 to tensor?x?xf32 // ... }这里?表示动态维而shape_sym属性把动态维和具体符号名绑定起来。这样 IR 里的每个 value 都自带一份完整的“shape 表达式”信息编译器 pass 可以做统一处理。同时还要在 IR 里维护一个约束集合例如#shape_constraint #shape.constraint N 1, N 64, L 8, L 512, N * L % 8 0 我把这个约束集合挂在包含 dynamic shape 的 func 或 region 上所有涉及 shape 的 pass 都先读取它。这不是锦上添花而是动态 shape 支持机制的核心光有表达式没有约束优化 pass 依然不敢做任何假设。3.3 约束集合是动态 Shape 的“隐藏裁判”很多框架实现动态 shape 时只做到“表达”没做到“约束”。结果就是编译期能画出 shape 表达式但无法回答“这个维度能不能被 32 整除”“这两个维度是不是一定相等”这类问题。没有约束就无法安全生成高效的向量化 kernel也无法做内存复用。约束集合需要支持几种常见关系相等关系batch input_0.dim0即两个符号代表同一个值可以 union-find 合并。区间关系1 seq_len 1024提供 range。有 range 之后编译器可以估算最大内存准备 buffer pool。对齐关系seq_len % 8 0这个信息对 SIMD 优化至关重要。算术关系full_len seq_len 2 * pad_len用于推导 reshape/concat 后的维度。约束集合的维护不是一次性完成的而是随着 shape inference pass 的处理不断叠加。每个 op 的 shape function 在生成输出 shape 时也会生成一些新的约束条件。编译器要做的是把这些约束收编进一个全局 solver/checker在涉及到分支判断时向它提问。我在工程实践中还发现一个窍门约束集合要同时支持“编译期可满足性检查”和“运行期校验”。编译期我们只需要知道某个假设是否有可能成立比如“如果 seq_len 能被 8 整除就可以走 vectorized kernel”这时候 solver 做一些区间推理就够了。运行期则要真正校验实际 shape 是否满足这些约束如果不满足就 fallback 到通用路径。4. 动态 Shape 支持机制在编译管线中的落地4.1 基于形状表达式的 shape inference动态 shape 的 shape inference 不能再用简单的常量传播每个算子都要有一个“符号 shape 推导规则”。举个例子矩阵乘的推导规则def infer_matmul(a: TensorMetadata, b: TensorMetadata) - TensorMetadata: # a: [M, K], b: [K, N] m a.shape[0] k_a a.shape[1] k_b b.shape[0] n b.shape[1] assert_symbolic_eq(k_a, k_b, matmul contract dim mismatch) return TensorMetadata( shape(m, n), dtypeinfer_matmul_dtype(a.dtype, b.dtype), ) def assert_symbolic_eq(dim_a, dim_b, message): # 如果两个维度都是符号加入相等约束 # 如果一个是常量另一个是符号尝试用区间约束证明。 solver.add_eq(dim_a, dim_b)注意这里的assert_symbolic_eq不能像静态 shape 那样直接在编译期报错。动态 shape 下两个符号维度是否相等可能要留到运行期才能完全确认。所以它实际上是“记录约束 运行期检查”而不是“编译期死磕”。真实项目里每个算子都要写类似的 shape function。Pointwise 算子最简单输出 shape 是对齐后的输入 shapeReduce、Concat、Reshape 各自有不同的推导逻辑。Complicated 的是 Reshape输入的总元素数是一个符号表达式输出 shape 必须能让这些元素数保持一致编译器只能生成一个运行期 assertion然后继续往下。4.2 形状特化bucketing 和 guard纯符号推导解决了“正确性”却解决不了“性能”。动态 shape 下如果所有算子都按最通用实现来编译性能基本没法看。所以在实际部署里更常见的策略是“形状特化”也就是把动态 shape 值域划分成若干桶bucket每个桶单独编译一套静态 shape 的变体。桶的划分一般结合线上 shape 分布。比如动态维度常见取值划分 bucketbatch size1, 2, 4, 8, 16, 32, 641, 4, 8, 16, 32, 64seq_len32~102464, 128, 256, 512, 1024图像 HxW多个分辨率320x320, 416x416, 512x512, 768x1024划分原则是桶的数量不要太多一般控制在 20~50 之间每个桶内的 shape 差异对性能影响不大极端 shape 不单独开桶走 fallback。编译期IR 会被复制成多个变体每个变体的动态符号被绑定到一个具体整数。这个绑定发生在 shape inference 之后、所有静态优化之前。每个变体前面插入一个 guard 节点运行时判断输入 shape 是否匹配当前变体。我再强调一遍这个 guard 是在运行时执行的不是编译期比较所以速度必须快通常用简单的整数比较就能完成。4.3 静态子图动态子图混合执行如果整张图所有地方都动态那编译优化空间就很小。实际模型通常只有一部分 shape 是动态的比如输入层和某些动态控制流其他大部分算子都是形状确定的。所以我会把计算图划分成静态子图和动态子图。划分规则很直接一个动态 shape 算子产生的 tensor如果被后续所有算子使用而后续算子又要求 shape 精确匹配那这些算子都要算作动态子图。如果某个动态 tensor 经过 padding、clip 或者特化后 shape 变成常量那此后的算子可以划回静态子图。边界上可能需要插入 layout conversion 或 memory copy成本要算在动态子图里。这种混合执行的好处是静态子图可以继续享受算子融合、常量折叠、内存复用动态子图则单独走动态 kernel 和动态分配。我在实际项目里见过一个模型动态输入进来后先做一轮pad_to_fixed, 之后几十个算子全部恢复成静态 shape速度几乎追平纯静态编译。4.4 内存计划和动态缓冲池动态 shape 对内存规划的影响最大。静态 shape 时arena 可以一次性把整张图的中间张量排布好动态 shape 时每个中间张量 size 都是运行期才知道传统 arena 方案失效。我的处理办法是两层分配策略静态区编译期能确定 size 的 tensor继续走静态 arena这部分内存可以复用。动态区size 只有在运行期确定的 tensor走一块独立的动态 buffer pool。每个 bucket 会预分配一块最大 size 的 buffer实际 shape 小于等于 bucket 上限时反复使用。动态区的复用策略比静态区宽松一些因为同 bucket 内的 tensor size 不一定完全相等但 buffer 容量是固定的可以复用。加上内存池扩容机制避免每次都走系统 malloc。另外还要注意对齐问题。动态 shape 取到的值经常是 3、5、17 这种不是 16 的倍数直接分配会导致 SIMD 访存效率下降。所以我一般会把动态维度向上对齐到 16 或 32计算时用 mask 或者 pad 处理边界。这个对齐信息也要写进 TensorMetadata不然运行时不记得 buffer 的有效长度和分配长度之间的差距。4.5 运行时的形状 dispatch最后落到运行时shape dispatch 逻辑其实是一个match_shape(inputs) - executable的过程。我会为每个编译变体生成一个 shape descriptor运行时用输入的 shape 去匹配。匹配顺序一般是完全相等匹配命中则直接执行。如果多个变体都能覆盖当前 shape选最近的一个 bucket避免频繁重编译。如果都匹配不上走 fallback解释执行或者调用一个通用动态 kernel。对于一些高频出现的未匹配 shape可以触发 on-the-fly JIT 编译编译结果缓存起来。运行时 dispatch 的开销必须控制住。形状比较不要每次都用 tuple 哈希最好把 shape 值打包成几个 64 位整数快速比较。否则动态 shape 的收益会被 dispatch 本身吃掉。5. 实操中的坑与排查实录5.1 符号碰撞所有输出 shape 突然变错动态 shape 最常见的问题是约束集合里两个本该独立的符号被错误等价了。比如模型有两个动态输入一个 batch size一个 beam size如果 shape function 里粗心地把它们 merge 到同一个符号那后面所有 shape 推导都会带上这个错误假设。我们当时排查了一个很诡异的 bug模型在 batch16、beam4 时正确batch8、beam4 时输出 shape 完全乱掉。最后打印符号绑定表才发现编译器把batch * beam当成一个新的符号又和某个中间维度的batch符号 union 了导致约束关系错误。排查手段就是“符号绑定表”。在 shape inference 之后打印每个符号对应的实际维度来源、区间约束、相等类。维护起来不复杂但对定位这类问题帮助极大。5.2 编译时间爆炸动态 shape bucketing 很容易产生笛卡尔积式的编译变体。比如 batch 有 6 个桶seq_len 有 6 个桶图像尺寸有 5 个桶三者的组合就是 180 个变体每个变体都要跑一遍完整编译优化几个小时都编不完。我最后的解法是用线上 shape 分布的 top-K 和覆盖率来决定桶而不是穷举所有组合。提供“编译预算”参数比如最多编译 32 个变体超过后直接使用通用 kernel。对模板特化按需 JIT而不是离线全量编译。在 kernel 层面尽量把 shape 作为运行时参数而不是编译期模板参数这样同一个 kernel 可以服务多个桶大幅减少编译次数。5.3 动态 shape 下算不对一个 reshape 引发的血案动态 shape 下最容易出错的是 reshape。静态 shape 时元素总数直接算好reshape 不可能非法。动态 shape 时[N, 128]reshape 成[N // 2, -1]看着没问题但如果N是奇数运行期就会崩。更隐蔽的是某些实现里-1推导会得到小数比如N * 128 / 24如果不整除结果完全错误。处理方式必须是编译期生成动态断言运行期检查整除性。不要尝试在编译期把所有情况证明完。我给 shape function 加了一个检查点所有 reshape 操作都会生成assert(total_elements % new_shape_prod 0)这样错误能提前暴露在推理阶段而不是等到 kernel 越界。5.4 调试技巧小结我把这几年常用的动态 shape 调试手段整理成一张速查表症状排查点快速修复输出 shape 莫名其妙变大检查符号绑定表看是否有符号被错误 union严格区分每个输入的符号名性能时好时坏检查 shape dispatch 是否频繁触发 JIT/fallback增加 hot shape 的 bucket 覆盖率内存持续上涨动态 buffer pool 没有按 bucket 回收加入空闲队列和最大容量限制偶发 shape mismatch检查 reshape 或 concat 的动态断言是否齐全补运行期 assert提前fail编译时间过长组合爆炸使用线上 top-K 桶限制变体数量结果数值错误但 shape 正确查看 strides 和 memory format 是否在动态路径里被忽略保留并检查非连续 tensor 标记还有一个通用调试开关我建议所有动态 shape 框架都支持打开后每个算子的输入和输出会打印两套 shape一套是编译期的符号表达式一套是运行期的真实 shape。两者不一致时直接报 warning。这个开关几乎能定位 90% 的动态 shape 问题。6. 几句题外话动态 Shape 支持的分寸感最后聊点个人体会。动态 Shape 支持机制不是越复杂越好。如果你的场景只是 batch size 变一变直接 padding 到固定 batch 或者做几个 bucket就能解决绝大多数问题。符号 shape 表达式和约束求解器这套东西工程成本不低如果不是要支持那种 shape 变化非常自由的模型不建议一上来就全上。我踩过的坑是一开始想着“做大而全的动态 shape 系统”结果编译器 pass 越写越多约束 solver 性能跟不上最后反而拖慢了上线进度。后来改成“静态为主、动态度量、bucket 兜底”的务实方案效果反而好很多。真正要把动态 Shape 做好元数据定义层的设计是关键。shape 表达式、约束集合、运行期 guard 这三样东西是整套机制的骨架。先把它们在 IR 里表达清楚后续任何 pass 都好做如果这一步偷懒后面只会越补越乱。做计算平台优化很多时候拼的不是技巧而是对“元数据契约”的尊重。动态 Shape 只是其中一个缩影。希望这篇内容对正在被 shape 困恼的你有帮助也欢迎在实践中多尝试不同的约束表示方式找到最适合你业务形态的那一套。