模型在 PyTorch 里跑得好好的,一转成 ONNX 就翻脸——喂进去一张 1080p 的图,直接甩你一行Got: (1, 3, 1080, 1920), Expected: (1, 3, 640, 640)。这种场面我遇到过太多次了,尤其是做车牌识别、人物抠图这类输入尺寸天生就千变万化的场景。ONNX 的动态输入和动态输出,本质上就是给张量的每一个维度两个身份选项:要么钉死成一个常数,要么写成一个符号名,让它在每次推理时跟着你喂进去的数据自己变。
这篇东西不打算讲 ONNX 的基础语法,那些文档里都有。我要聊的是从导出、验证、推理到跨平台部署这一整条链路上,动态 shape 到底会在哪些地方给你使绊子,以及每个绊子背后真正的成因。如果你正卡在"同一个模型换张图就报错"、"转 TensorRT 后只能跑一个尺寸"、"量化完精度崩了"这类问题上,这篇的排查路径应该能省你不少时间。文中给的代码都是可以直接抄去改的,我也会把参数为什么要这么填讲清楚。
1. 动态维度在 ONNX 图里到底是怎么表达的
1.1 一个维度只有三种状态,别把它想复杂
ONNX 里描述形状用的是TensorShapeProto,里面是一串Dimension。每个Dimension只有三种可能的状态,理解这三种状态,后面所有的报错你都能自己定位。
第一种是定值维,dim_value被赋了一个具体的整数,比如dim_value = 1。这种维度在推理时不许变,你喂进去的形状对不上就直接抛异常。
第二种是符号维,dim_param被赋了一个字符串,比如dim_param = "batch"或者dim_param = "height"。这个字符串本身没有数值含义,它只是一个"占位标签"。关键在于:同一个标签在同一个图里多次出现时,ONNX Runtime 会认为它们是同一个值。这一点非常重要,也是很多诡异的 shape 报错的根源——你把输入高度标成了height,而图内某个中间张量恰好也用了同一个标签,ORT 就会尝试去满足这个约束,结果可能跟你预期完全不一样。
第三种是完全空,既没有dim_value也没有dim_param。这表示"这里有一个维度,但大小未知"。onnx.shape_inference在推不出来的时候就会留下这种空维度。
提示:不少人以为只要把输入标成符号维就万事大吉了,其实图内部的中间张量如果被推成了
dim_value,动态性会在中间某一层被掐断。验证的时候不能只看输入输出,要打印全部中间张量的形状。
1.2dim_value和dim_param是互斥的,重复设置会静默失效
手写代码去修形状的时候,最容易踩的一个坑是这样写的:
import onnx model = onnx.load("model.onnx") dim = model.graph.input[0].type.tensor_type.shape.dim[0] dim.dim_value = 1 # 先设了定值 dim.dim_param = "batch" # 又想改成符号维 onnx.save(model, "patched.onnx")这段代码不会报错,onnx.checker大概率也能过(取决于版本),但加载进 ORT 之后你会发现动态性根本没生效,或者报一个Invalid dimension之类的错。原因是 protobuf 的 oneof 语义:dim_value和dim_param属于同一个 oneof 字段,后写的那个会覆盖前一个,但如果原来的值是默认值 0,有时候序列化出来又是另一回事。稳妥的写法永远是先彻底清空再设:
from onnx import TensorProto def set_dim(dim, name=None, value=None): # 清空 oneof 里的两个字段 dim.ClearField("dim_value") dim.ClearField("dim_param") if name is not None: dim.dim_param = name elif value is not None: dim.dim_value = value for inp in model.graph.input: dims = inp.type.tensor_type.shape.dim if len(dims) == 4: set_dim(dims[0], name="batch") set_dim(dims[2], name="height") set_dim(dims[3], name="width")ClearField是 protobuf 的标准 API,对 oneof 字段用它是唯一可靠的做法。这个细节几乎没有文档会提,但它坑过的人不在少数。
1.3 稀疏维和稠密维的区别,别搞混
ONNX 的类型系统里,TypeProto有一个 oneof,可以是tensor_type,也可以是sparse_tensor_type、sequence_type、map_type、optional_type。你在改形状时必须确认自己改的是tensor_type,否则改半天没反应。写个带校验的工具函数会省事很多:
def is_tensor(vi): return vi.type.HasField("tensor_type") def get_shape(vi): if not is_tensor(vi): return None tt = vi.type.tensor_type if not tt.HasField("shape"): return None out = [] for d in tt.shape.dim: if d.HasField("dim_value"): out.append(d.dim_value) elif d.HasField("dim_param"): out.append(d.dim_param) else: out.append(None) # 未知维 return out这里None表示未知维,注意它和符号维不是一回事。未知维在图优化阶段可能被折叠成任何值,而符号维有明确的命名,ORT 的 free dimension override 只能作用于符号维。这个区别在后面讲 ORT 性能调优时还会用到。
2. 导出阶段就把动态轴钉死:从 dummy input 到符号传播
2.1torch.onnx.export里的dynamic_axes到底改了什么东西
很多人对dynamic_axes有一个误解,以为它会"让整个模型支持动态形状"。它实际做的事情只有一件:修改图输入和图输出的形状声明。真正让图内部算子的形状推导变成符号式的,是 PyTorch 导出器在追踪过程中做的符号传播。
一个典型写法是这样的:
import torch model.eval() dummy = torch.randn(1, 3, 640, 640) torch.onnx.export( model, dummy, "det.onnx", input_names=["images"], output_names=["preds"], dynamic_axes={ "images": {0: "batch", 2: "height", 3: "width"}, "preds": {0: "batch"}, }, opset_version=13, do_constant_folding=True, )注意dynamic_axes的键是你在input_names/output_names里起的名字,不是 PyTorch 里参数的变量名。名字写错了不会报错,只是那一条动态声明被静默忽略——这是最隐蔽的坑之一。写完导出脚本的第一件事,应该是立刻用 ORT 把输入形状打出来看。
PyTorch 2.x 之后多了一条dynamo=True的导出路径,形状控制改成了dynamic_shapes,语义上用Dim对象描述,表达能力更强(比如可以表达"高度等于宽度"这种约束)。如果你的环境是较新的版本,两条路都值得试一下,老路径对某些控制流更宽容,新路径对动态 shape 的支持更彻底。具体用哪个,还是看你模型里有没有 Python 风格的if/for。
2.2 名字对不上会出什么乱子
如果模型forward返回的是一个 tuple,而output_names只给了一个名字,导出器会把剩下的输出自动命名成12、13这种数字名字。后面用 C++ 或 Java 去取输出时按名字拿,就会拿到 null。反过来,如果output_names给的数量比实际输出多,导出会直接失败。
一个稳妥的做法是在导出后立刻做一次一致性检查:
import onnxruntime as ort sess = ort.InferenceSession("det.onnx", providers=["CPUExecutionProvider"]) for i in sess.get_inputs(): print("IN ", i.name, i.shape, i.type) for o in sess.get_outputs(): print("OUT", o.name, o.shape, o.type)打印出来的shape里,字符串就是符号维,整数就是定值维。看到['batch', 3, 'height', 'width']才说明动态轴真的生效了;如果打印出来是[1, 3, 640, 640],那就是没写对名字,或者被后面的优化步骤折叠回去了。
2.3 实测验证:同一份 session 跑两种尺寸
验证动态性最直接的办法就是用同一个 session 连续跑两个不同尺寸:
import numpy as np for hw in [(640, 640), (1080, 1920), (384, 1280)]: x = np.random.rand(1, 3, *hw).astype(np.float32) y = sess.run(None, {"images": x}) print(hw, "->", [t.shape for t in y])这里有个实操细节值得说:每次换尺寸,ORT 会重新做一次内存规划和 kernel 选择。第一次跑新尺寸会比较慢,后面同尺寸就快了。所以测性能的时候一定要先 warmup 几轮再计时,否则你会得到"动态 shape 比静态慢十倍"的错误结论。
2.4 已经导出成静态的了,还能不能救
能救,但要分清"改声明"和"改计算图"两件事。onnx.load之后直接改graph.input的dim_param,改的是声明;而图内部的Reshape、Resize、Squeeze这些算子的 shape 输入往往是常量,它们不受声明影响,会继续把你锁死在原来的尺寸上。
所以补救流程通常是三步:先改输入输出声明,再跑onnx.shape_inference.infer_shapes看符号能传播到哪一层,最后用onnx-simplifier做一次常量化简和冗余清理:
python -m onnxsim det_static.onnx det_dyn.onnx --dynamic-input-shape --overwrite-input-shape--dynamic-input-shape会把所有输入的批次维和空间维都设成动态。但如果 sim 之后你去跑一个大尺寸输入发现还是报错,那说明图内部有写死的常量在卡着,只能回到导出环节重新导一遍。这种情况下,改导出脚本比事后修补省事得多。
3. 动态输出比动态输入麻烦得多
3.1 输出为什么会变长:NMS、TopK 和变长解码
输入动态是"你告诉模型可以多大",输出动态是"模型告诉你它找到了多少个"。后者通常来自三类算子。
第一类是非极大值抑制(NMS)。检测模型输出的候选框数量本身是动态的,NMS 之后保留多少个完全取决于图像内容。ONNX 里对应的是NonMaxSuppression算子,它的输出是[num_selected, 3],第一维天然是符号维。
第二类是动态 TopK。语音识别、检索类模型里常见,选出前 k 个候选,k 可以动态也可以固定,但排序后的索引范围是动态的。
第三类是控制流。ONNX 支持If和Loop两个算子,能表达数据依赖的分支和循环。带Loop的图,输出形状往往要到运行时才知道。
3.2 动态输出最坑的地方:下游代码拿不到形状
静态模型里,你可以这么写:
out = sess.run(None, feeds)[0] boxes = out[0, :, :4] # 形状写死换成动态输出就不行了,因为第一维是变量。必须改成从运行时结果里读:
outs = sess.run(None, feeds) preds = outs[0] print("实际输出形状:", preds.shape) # 每次都可能不同C++ 和 Java 侧同理。Java 里result.get(0).getValue()拿到的是嵌套数组,长度必须从数组本身读,不能靠预先分配的固定 buffer。如果你在 Java 里看到ArrayIndexOutOfBoundsException或者结果被截断,八成就是这里写死了长度。
3.3 我更推荐的做法:把 NMS 挪出 ONNX 图
经过这么多年折腾,对检测类模型我有一个比较明确的偏好:尽量让 ONNX 图的输出形状固定,把 NMS 放到图外做。
YOLO 系列的 anchor-free 输出就是个很好的例子,模型主干输出[1, N, 85],N 只跟输入分辨率有关,跟图像内容无关,形状是完全可预测的。NMS 用 numpy、Java 或者 C++ 写一遍,几十行代码,性能比图内实现还好,因为你可以顺便做阈值裁剪和 top-k 截断。
这样做的好处是连锁的:ONNX 图变简单之后,量化更容易、转 TensorRT 更容易、跨平台一致性更好,Java 侧也不用处理变长数组。代价是你得自己写后处理,但这段代码写一次就够了。
如果确实必须在图内做 NMS,那就得接受输出动态,同时注意NonMaxSuppression在不同 opset 版本里的输入参数有差异,center_point_box这个属性的默认值在 v11 和 v13 之间是不一致的,跨版本转换时经常对不上。
4. 推理阶段的实战坑:打包策略、内存和预分配
4.1 动态 batch 不等于无脑加大 batch
批处理能提升吞吐,但动态 batch 有一个绕不开的问题:同一个 batch 里的样本必须 padding 到相同尺寸。如果一批里有 640x640 的图和 1920x1080 的图,padding 到 1920x1080 之后,前者有超过 60% 的像素是无效的。算力全浪费在零上了。
我一般用三种策略,按场景选:
| 策略 | 适用场景 | 优点 | 代价 |
|---|---|---|---|
| batch=1 + 动态尺寸 | 实时单路视频流 | 无 padding 浪费,延迟最低 | 吞吐低,GPU 利用率不高 |
| 定尺分桶 | 离线批量处理 | padding 浪费可控,吞吐高 | 需要先统计尺寸分布 |
| 动态尺寸 + batch | 尺寸接近的批量任务 | 灵活 | 排序开销大,实现复杂 |
"定尺分桶"是我用得最多的。做法是先跑一遍数据集,统计所有图片的长宽,聚成 4 到 8 个桶(比如 640x640、960x544、1280x736、1920x1088),同一批只放同一个桶的图。这样每个桶里可以做静态 shape,ONNX Runtime 的图优化和内存复用都能吃到,实测吞吐比纯动态 shape 高出不少。
4.2 动态 shape 会让 ORT 的内存池失效
ORT 默认用一个 arena 分配器来管理内存,同一个 session 反复跑相同形状时,内存是复用的,几乎不产生新的分配开销。但形状一变,arena 里的块就对不上,需要重新申请。极端情况下(比如视频流里每一帧尺寸都不同),你会看到内存占用持续抖动,甚至因为碎片而增长。
处理方式有两个。一是限制形状的种类数——比如把输入尺寸都对齐到 32 的倍数再送入,这样尺寸组合会收敛很多。二是用IOBinding自己管理输入输出内存:
import onnxruntime as ort import numpy as np so = ort.SessionOptions() so.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL so.enable_mem_pattern = False # 形状多变时关掉内存模式 sess = ort.InferenceSession("det.onnx", so, providers=["CPUExecutionProvider"]) io = sess.io_binding() x = np.zeros((1, 3, 640, 640), dtype=np.float32) io.bind_cpu_input("images", x) io.bind_output("preds") sess.run_with_iobinding(io) out = io.get_outputs()[0].numpy()enable_mem_pattern在形状固定时开着能省内存,形状多变时反而会拖慢首次推理,值得根据场景测一测。
4.3 用 free dimension override 把符号维焊死
如果你导出的模型是动态的,但实际部署时尺寸是固定的,可以在 session 层面把符号维替换成常数。这样既能复用同一份动态模型,又能拿到静态图的优化效果:
so = ort.SessionOptions() so.add_free_dimension_override_by_name("height", 640) so.add_free_dimension_override_by_name("width", 640) sess = ort.InferenceSession("det.onnx", so, providers=["CPUExecutionProvider"])这个 API 只对dim_param生效,对未知维(既没有 value 也没有 param)无效——这也是我在第 1.3 节强调未知维和符号维要分清的原因。
注意:override 之后,模型就只接受这一个尺寸了,喂别的尺寸会报形状不匹配。适合"模型动态、部署固定"这一种场景,别在需要真正多尺寸的场景里用。
5. 跨平台部署时的动态 shape 处理
5.1 转 TensorRT:profile 的三段式必须配合业务峰值
trtexec构建带动态维度的 engine 时,要指定minShapes/optShapes/maxShapes:
trtexec --onnx=yolo.onnx --saveEngine=yolo_fp16.engine --fp16 \ --minShapes=images:1x3x320x320 \ --optShapes=images:1x3x640x640 \ --maxShapes=images:1x3x1280x1280三个形状的含义要搞清楚:optShapes是性能拐点,TensorRT 会按它来选 kernel 和调优;maxShapes决定显存占用的上限,因为 TensorRT 是按最大形状预分配 workspace 的。很多人只填了 max 填得很大,结果显存爆了,其实是可以把 max 压到业务真实峰值上线的。
另外,engine 一旦构建完成,形状范围就固定了。如果实际输入超出了 profile 范围,TensorRT 不会自动缩放,会直接报错。对于分辨率变化很大的场景,我的做法是构建两三个 profile 或者两个 engine,按输入尺寸分流。
5.2 转 RKNN:动态 shape 基本要走固化这条路
把 ONNX 转到 RKNN 部署在边缘设备上时,动态 shape 的支持非常有限,尤其是 int8 量化路径。原因是 int8 静态量化需要在构建阶段用校准数据集跑一遍,统计每一层激活值的分布来确定量化参数。如果输入尺寸每次都不同,激活分布没法收敛,量化误差会非常不可控。
所以流程上要倒过来:先确定部署尺寸,把 ONNX 图改成固定 shape,再做量化。如果业务上确实需要多种尺寸,就在同一份权重上导出多个固定尺寸的 ONNX,分别转成多个 RKNN 模型,运行时按需选择。听起来笨,但这是目前最稳的做法。
校准集也有讲究:不要随便抓几张图就用,最好覆盖实际业务里的各种光照、天气、角度。我一般取 200 到 500 张,覆盖典型场景,太多也没用,反而拉长构建时间。
5.3 int8 量化与动态 shape 的组合坑
onnxruntime.quantization提供了两条路。quantize_dynamic只对矩阵乘法类的算子生效(MatMul、Gemm、Attention 这类 Transformer 结构),卷积网络基本吃不到收益。CNN 检测模型要走quantize_static,需要提供校准数据读取器。
这两条路和动态 shape 的关系不太一样。动态量化对形状不敏感,因为它是权重的离线量化,激活值在运行时动态算 scale,所以动态 shape 模型也能跑。但静态量化对形状敏感,校准过程中的激活统计是跟输入尺寸绑定的,尺寸一变,量化 scale 就失配了。
实践中我的经验是:如果模型必须保持动态 shape,就不要做静态 int8 量化;如果要做 int8,就先把形状固化下来。硬要两者兼得,只能自己在图里插QuantizeLinear/DequantizeLinear节点做 QDQ 格式的量化,控制粒度更细但工作量翻倍。
5.4 Java 端 onnxruntime 处理动态尺寸的几个要点
Java 生态里做车牌识别、OCR 检测这类活儿,ONNX Runtime 的 Java API 是主流选择。几个和动态 shape 直接相关的点:
OrtEnvironment env = OrtEnvironment.getEnvironment(); OrtSession.SessionOptions opts = new OrtSession.SessionOptions(); opts.setOptimizationLevel(OrtSession.SessionOptions.OptLevel.ALL_OPT); OrtSession session = env.createSession("det.onnx", opts); int h = 640, w = 1280; long[] shape = new long[]{1, 3, h, w}; FloatBuffer buf = FloatBuffer.wrap(nchwArray); OnnxTensor input = OnnxTensor.createTensor(env, buf, shape); Map<String, OnnxTensor> feeds = new HashMap<>(); feeds.put("images", input); try (OrtSession.Result res = session.run(feeds)) { float[][][][] out = (float[][][][]) res.get(0).getValue(); // 注意:out[0].length 是运行时才知道的,不要写死 }第一个要点是shape数组的乘积必须和FloatBuffer的长度严格相等,否则会抛异常,而且异常信息往往不直观。第二个要点是输出数组的维度长度要动态读,res.get(0).getValue()返回的是嵌套数组,第一维长度就是实际输出数量。第三个要点是OrtSession是线程安全的,但OnnxTensor不是,多线程推理时要每个线程各自创建输入张量。
至于车牌识别里常见的 OCR 检测网络,它的输入高度通常固定(比如 32 或 48),宽度是动态的,而且宽度必须是 32 的整数倍,因为检测网络里有 5 次下采样,宽度不整除会导致最后一层的特征图尺寸对不上,报错位置还特别深,很难定位。我的做法是在送进模型之前统一做一次对齐:
int alignedW = (int) (Math.ceil(rawW / 32.0) * 32);这个 32 是从网络的下采样倍率来的,换个网络可能要改成 8 或 16,看结构定。
6. 一套可以直接照着走的排查顺序
遇到动态 shape 相关的问题,我现在的排查顺序基本固定下来了,从外到内一层层剥:
| 现象 | 大概率原因 | 排查动作 |
|---|---|---|
| 换尺寸就报形状不匹配 | 导出时动态轴没生效 | 打印 session 输入 shape,看是不是字符串 |
| 输入动态了但中间层对不上 | Reshape/Resize 的 shape 输入是常量 | 跑 shape_inference,检查中间张量 |
| 输出形状每次都不一样导致下游崩 | NMS/TopK 变长输出 | 打出运行时 shape,改下游取值逻辑 |
| ORT 内存持续增长 | 形状种类太多,arena 碎片 | 关掉 mem_pattern,或做尺寸分桶 |
| 转 TensorRT 后只能跑一个尺寸 | profile 没配或配得不对 | 检查 min/opt/max 三段 |
| 量化后精度暴跌 | 动态 shape 下做了静态量化 | 固化形状重新校准 |
具体操作上,我会按这个顺序走一遍。第一步,用 ORT 加载模型,把输入输出的名字和形状全打出来,确认符号维的存在。第二步,用两个差异很大的尺寸各跑一次,看是否都成功,顺便看输出形状的变化。第三步,如果失败,用onnx.shape_inference.infer_shapes生成一个带形状信息的模型,然后遍历所有value_info找那些被推成常量的中间张量,那个位置就是动态性的断点。第四步,回到导出脚本,针对性地把那个断点上游的算子换掉或者包一层。
第三步是很多人会跳过的一步,但它其实是最高效的定位手段。写个小脚本遍历一遍,比在导出脚本里反复试参数快得多:
import onnx m = onnx.load("det.onnx") m = onnx.shape_inference.infer_shapes(m) for vi in m.graph.value_info: dims = [] for d in vi.type.tensor_type.shape.dim: if d.HasField("dim_param"): dims.append(d.dim_param) elif d.HasField("dim_value"): dims.append(d.dim_value) else: dims.append("?") # 打印那些本该动态却变成了常量的张量 if 0 in dims or 1 in dims: print(vi.name, dims)跑完之后你会看到一批张量,其中像Reshape的第二个输入、Resize的sizes输入这些,如果它们是常量张量,就是断点的来源。常见的修法是把尺寸计算挪到图里用Shape/Gather/Concat动态算出来,而不是在导出时用 Python 的常量算好。这个改法在 torch 侧就是把x.view(x.size(0), -1)这种写法改得对 shape 算子更友好,或者干脆用torch.onnx.export的dynamo=True路径让它自己处理。
最后再分享一个我踩过好几次的经验:导出完的模型一定要拿两三个差异极大的尺寸各跑一遍再做交付,别只用导出时那个 dummy 尺寸验证。我就遇到过一次,导出脚本里动态轴配得没问题,用 640 验证也过,结果线上第一个 1080p 的请求就挂了——原因是模型里有一个interpolate用了固定的scale_factor,图内部的尺寸在某一层被常量折叠成了导出时的值。这种问题只有用不同的尺寸实测才能暴露出来,看代码是看不出来的。现在我的习惯是把尺寸实测写进 CI,模型一更新就自动跑一遍多尺寸校验,省得线上再踩一次。