1. 为什么导出的 ONNX 总是被“焊死”成固定维度
如果你用 PyTorch 训练完一个文本分类或序列标注模型,兴冲冲地torch.onnx.export出来,然后拿一条长度 128 的句子推理没问题,换成长度 64 的句子直接报 shape mismatch——大概率不是模型写错了,而是导出时输入输出维度被 dummy input 的形状“焊死”了。
ONNX 的图结构里,每个张量的 shape 默认是静态的。你喂进去的dummy_input是[1, 128],导出的模型就认为第一维永远是 1、第二维永远是 128。可变 batch、变长序列这些部署时最常见的需求,全被这个默认行为挡在门外。
这篇要解决的就是这件事:怎么在torch.onnx.export阶段就用dynamic_axes把维度声明成动态的,让一次导出适配多种 shape;导出后怎么用 onnxruntime 实际跑不同尺寸的输入来验证;以及当dynamic_axes不生效、中间节点仍然是死维度时,怎么排查和补救。
适合的人群:做 NLP 序列模型部署、需要可变 batch 推理、或者被 ONNX 固定维度坑过的同学。下面所有配置和命令都可以直接复制改路径使用。
2. 前置准备:TaoToken 配置骨架与依赖环境
在动手改导出脚本之前,先把环境和访问凭证理顺。我习惯把模型导出、验证脚本放在同一个工程里,用 TaoToken 统一管理模型调用和 API 访问的配置,这样后面如果要接在线推理或做对比验证,不用再单独折腾一套 key。
TaoToken 的官网入口是 https://taotoken.net/?utm_source=taotoken_aicg_blog_end&utm_medium=csdn&utm_campaign=rewrite&utm_content= ,API 基址是 https://taotoken.net/api (这个不加 UTM)。如果你只是本地导出 ONNX、用 onnxruntime 验证,其实不依赖任何在线服务;但如果你想把导出后的模型接到一个统一的推理服务里做端到端验证,或者需要调用模型对话能力做结果比对,那提前把 key 配好会省事很多。
依赖安装这块,核心就三个包:
pip install torch onnx onnxruntime numpy版本上不用太纠结,PyTorch 1.10 以上、onnx 1.12 以上、onnxruntime 1.14 以上基本都支持本文的写法。如果你用的是 GPU 环境,onnxruntime 换成onnxruntime-gpu即可,验证脚本代码不用改。
配置骨架我一般写成这样,放在config.py里:
# config.py import os TAOTOKEN_API_BASE = os.getenv("TAOTOKEN_API_BASE", "https://taotoken.net/api") TAOTOKEN_API_KEY = os.getenv("TAOTOKEN_API_KEY", "") ONNX_MODEL_PATH = "exports/model_dynamic.onnx" DUMMY_BATCH = 1 DUMMY_SEQ_LEN = 128把 key 放在环境变量里,不要硬编码进脚本。需要生成或管理 key 的话,控制台在 https://taotoken.net/console?utm_source=taotoken_aicg_blog_end&utm_content=console&utm_campaign=rewrite ,API Keys 管理页在 https://taotoken.net/api-keys?utm_source=taotoken_aicg_blog_end&utm_content=api-keys&utm_campaign=rewrite 。接入文档在 https://taotoken.net/doc?utm_source=taotoken_aicg_blog_end&utm_content=doc&utm_campaign=rewrite ,遇到参数不确定的时候翻一下比猜快。
注意:本文的 ONNX 导出和验证完全在本地完成,TaoToken 的作用是给你一个统一的配置入口和后续在线验证的通道,不是导出流程的必需依赖。别把它当成“导出工具”来理解。
3. 可复制的 dynamic_axes 配置与 export 参数骨架
这是全文最核心的一段。先看一个典型的、会导出成固定维度的错误写法:
import torch class TinySeqModel(torch.nn.Module): def __init__(self, vocab=1000, hidden=64, num_class=2): super().__init__() self.emb = torch.nn.Embedding(vocab, hidden) self.lstm = torch.nn.LSTM(hidden, hidden, batch_first=True) self.fc = torch.nn.Linear(hidden, num_class) def forward(self, input_ids): x = self.emb(input_ids) x, _ = self.lstm(x) x = x[:, -1, :] return self.fc(x) model = TinySeqModel().eval() dummy_input = torch.randint(0, 1000, (1, 128)) torch.onnx.export( model, dummy_input, "model_fixed.onnx", input_names=["input_ids"], output_names=["logits"], opset_version=13, )这段跑完,model_fixed.onnx的输入就是[1, 128],batch 和序列长度都动不了。正确做法是加dynamic_axes:
import torch model = TinySeqModel().eval() dummy_input = torch.randint(0, 1000, (1, 128)) dynamic_axes = { "input_ids": {0: "batch_size", 1: "seq_len"}, "logits": {0: "batch_size"}, } torch.onnx.export( model, dummy_input, "model_dynamic.onnx", input_names=["input_ids"], output_names=["logits"], dynamic_axes=dynamic_axes, opset_version=13, do_constant_folding=True, )dynamic_axes的语义是:对名为input_ids的输入,第 0 维命名为batch_size,第 1 维命名为seq_len;对名为logits的输出,第 0 维命名为batch_size。名字是自定义的字符串,只要不是纯数字,ONNX 就把它当作符号维度(symbolic dimension),而不是固定值。
几个容易踩的点,我列成表格对照:
| 参数 | 作用 | 常见错误 |
|---|---|---|
input_names | 给输入起名,dynamic_axes 靠这个名字索引 | 名字和 dynamic_axes 的 key 不一致,静默失效 |
output_names | 给输出起名 | 输出维度没声明,batch 仍然固定 |
dynamic_axes | 声明哪些维度可变 | 只声明输入不声明输出,或维度索引写错 |
opset_version | 算子集版本 | 太低不支持某些动态算子,建议 12 以上 |
do_constant_folding | 常量折叠优化 | 一般保持 True,但某些动态 shape 场景要关掉 |
如果你的模型有多个输入(比如input_ids+attention_mask),每个都要单独声明:
dynamic_axes = { "input_ids": {0: "batch_size", 1: "seq_len"}, "attention_mask": {0: "batch_size", 1: "seq_len"}, "logits": {0: "batch_size"}, }输出如果有多个,同理逐个写。维度索引从 0 开始,{0: "batch_size"}表示第 0 维动态。命名建议统一用batch_size、seq_len这种语义化名字,方便后面排查。
4. 用 onnxruntime 加载并跑通不同 shape 的验证动作
导出完不能只看文件生成了就完事,必须实际用不同 shape 跑一遍。下面这个验证脚本可以直接用:
import numpy as np import onnxruntime as ort sess = ort.InferenceSession("model_dynamic.onnx", providers=["CPUExecutionProvider"]) def run(shape): input_ids = np.random.randint(0, 1000, size=shape).astype(np.int64) outputs = sess.run(["logits"], {"input_ids": input_ids}) print(f"input shape={shape} -> output shape={outputs[0].shape}") run((1, 128)) run((4, 128)) run((1, 64)) run((8, 32))预期输出类似:
input shape=(1, 128) -> output shape=(1, 2) input shape=(4, 128) -> output shape=(4, 2) input shape=(1, 64) -> output shape=(1, 2) input shape=(8, 32) -> output shape=(8, 2)如果(4, 128)或(1, 64)报错,说明动态维度没生效。这时候先检查dynamic_axes的 key 是否和input_names完全一致,再检查维度索引有没有写反。
想更直观地看 ONNX 图里每个节点的 shape 信息,可以用:
import onnx model = onnx.load("model_dynamic.onnx") for inp in model.graph.input: dims = [d.dim_param or d.dim_value for d in inp.type.tensor_type.shape.dim] print("input:", inp.name, dims) for out in model.graph.output: dims = [d.dim_param or d.dim_value for d in out.type.tensor_type.shape.dim] print("output:", out.name, dims)正常应该打印出['batch_size', 'seq_len']这样的符号名,而不是[1, 128]这种数字。如果打印出来还是数字,说明导出时 dynamic_axes 根本没被识别。
5. 本篇常见错排查:dynamic_axes 不生效与中间节点死维度
错误一:dynamic_axes 的 key 和 input_names 对不上。这是最高频的坑。比如input_names=["input"]但 dynamic_axes 写的是{"input_ids": ...},ONNX 不会报错,直接忽略,导出结果还是固定维度。排查方法就是上面那段打印 input/output dims 的代码,看符号名在不在。
错误二:只改了输入没改输出。输入动态了,但输出logits的第 0 维还是固定 1。推理时 batch=4 输入能进去,输出却只有 1 行,后面接的逻辑全乱。输出维度一定要一起声明。
错误三:中间节点仍然是死维度。这就是 excerpt 里提到的情况——你改了 graph 的 input/output 维度,但网络内部某些节点的 shape 在导出时已经被常量折叠或算子推导固定住了。表现是:输入输出看着是动态的,但换个 shape 跑就报某个中间节点的维度不匹配。
排查这种问题,用 onnxruntime 的 verbose 日志:
import onnxruntime as ort so = ort.SessionOptions() so.log_severity_level = 1 sess = ort.InferenceSession("model_dynamic.onnx", so, providers=["CPUExecutionProvider"])日志里会指出哪个节点在哪个维度上失败。常见原因是模型里有view、reshape、squeeze这类对 shape 敏感的算子,写死了某个维度。解决办法是在 PyTorch 侧把这些操作改成动态友好的写法,比如用x.reshape(x.size(0), -1)而不是x.view(1, -1),或者用torch.nn.functional.adaptive_avg_pool1d替代固定窗口的池化。
错误四:opset 版本太低。某些动态 shape 相关的算子需要 opset 12 以上才支持。如果你用的是很老的 PyTorch,默认 opset 可能是 9 或 10,动态维度会出问题。显式指定opset_version=13或更高。
错误五:导出后手动改 dim_param 但没重跑验证。有人用onnx.load改dim_param再onnx.save,这招对简单的输入输出节点有效,但正如 excerpt 所说,中间节点的问题它解决不了。改完必须用第 4 节的脚本重新跑不同 shape,别只看文件保存成功。
6. 语义一致的收尾:把验证动作固化进你的导出流程
导出 ONNX 这件事,最怕的就是“导出了、没报错、上线才发现维度不对”。我的做法是把第 3 节的导出和第 4 节的验证合成一个脚本,导出后立刻跑一组不同 shape,全部通过才算成功。这样每次改模型结构或升级 PyTorch,都能第一时间发现动态维度有没有被破坏。
如果你后面要把这个模型接到统一的推理服务里做端到端测试,或者需要调用模型对话能力做输出比对,可以在验证脚本里通过 TaoToken 的 API 基址 https://taotoken.net/api 接入,key 从 https://taotoken.net/api-keys?utm_source=taotoken_aicg_blog_end&utm_content=api-keys&utm_campaign=rewrite 拿。模型对话入口在 https://taotoken.net/models?utm_source=taotoken_aicg_blog_end&utm_content=models&utm_campaign=rewrite ,接入文档在 https://taotoken.net/doc?utm_source=taotoken_aicg_blog_end&utm_content=doc&utm_campaign=rewrite 。如果你在做长期的编码或 Agent 类项目,需要稳定的调用额度,可以看 Coding Plan:https://taotoken.net/coding-plan?utm_source=taotoken_aicg_blog_end&utm_content=coding-plan&utm_campaign=rewrite 。
最后留一个我踩过的坑:dynamic_axes里的维度名不要用?这种单字符,虽然某些工具能识别,但 onnxruntime 在部分版本下对?的处理不一致,用batch_size、seq_len这种明确的名字最稳。导出脚本里加一行打印 input/output dims,比事后 debug 省太多时间。