如何识别TorchScript模型的输入形状?
TorchScript模型查询输入形状方法
答案是可以,你可以直接通过TorchScript模型文件获取输入形状相关信息,具体操作方法如下:
- 方法1:通过计算图
graph属性直接读取输入节点信息
运行以下代码即可直接打印模型所有输入的类型、形状信息:
输出结果中第一个输入一般是模型自身的import torch # 加载本地TorchScript模型文件 model = torch.jit.load("your_model_file.pt", map_location="cpu") # 遍历并打印计算图的所有输入节点 for input_node in model.graph.inputs(): print(input_node)self参数,后续的就是推理时需要传入的输入张量,形如Tensor(shape=[1, 3, 224, 224], dtype=Float32)的输出内容,就对应了输入的形状和数据类型要求。 - 方法2:通过
code属性查看前向逻辑的结构化信息
如果模型逻辑不复杂,还可以直接打印模型的伪代码查看输入参数定义:
输出内容会直接展示前向传播函数的入参定义,以及每一步的运算逻辑,输入形状信息会直接标注在入参的类型声明中。print(model.code) - 方法3:动态形状模型的补充验证
如果你查到的输入形状带有*标记(比如Tensor(shape=[*, 3, *, *])),说明模型支持动态尺寸输入,你可以结合模型所属的任务类型构造不同尺寸的Dummy输入验证合法输入范围:- CV类模型常规输入格式为NCHW(批次、通道、高度、宽度)
- NLP类模型常规输入格式为[批次大小, 序列长度]
注意事项
- 如果模型是通过torch.jit.trace方式导出的,导出时使用的输入形状是固定值,查询到的形状就是推理时的固定输入要求
- 如果模型是通过torch.jit.script方式导出的,大概率支持动态尺寸输入,查询结果会显示可变维度的
*标记
内容的提问来源于stack exchange,提问作者Manveru
相关产品推荐
相关产品推荐

