使用torch.onnx.export导出时能否直接添加模型描述及自定义元数据?
解决方案
完全可以在调用torch.onnx.export时直接添加模型描述和自定义元数据,无需二次读写ONNX文件,PyTorch 1.10及以上版本原生支持该能力,具体操作如下:
核心修改点
调用torch.onnx.export时新增两个参数即可:
doc_string:字符串类型,直接对应ONNX模型的description字段,也就是你当前通过model.doc_string设置的内容metadata_props:字典类型,可传入版本号、训练时间等任意自定义元数据,后续可从ONNX Runtime的模型元数据中直接读取
修改后可运行示例
import torch import torchvision from onnxruntime import InferenceSession # 获取示例模型 dummy_input = torch.randn(10, 3, 224, 224) _model = torchvision.models.alexnet(pretrained=True) input_names = [ "actual_input_1" ] + [ "learned_%d" % i for i in range(16) ] output_names = [ "output1" ] # 导出ONNX模型,直接添加描述和自定义元数据 torch.onnx.export( _model, dummy_input, "alexnet.onnx", verbose=False, input_names=input_names, output_names=output_names, strip_doc_string=False, # 直接设置模型描述 doc_string="my_description", # 添加自定义元数据,可根据需求扩展任意键值对 metadata_props={ "internal_version": "v1.2.0", "train_dataset_version": "v202405", "pytorch_version": torch.__version__ } ) # 加载导出的ONNX模型执行推理 sess = InferenceSession('alexnet.onnx') meta = sess.get_modelmeta() # 查看描述字段 print(meta.description) # 输出 my_description # 查看自定义元数据 print(meta.custom_metadata_map["internal_version"]) # 输出 v1.2.0
兼容说明
如果你使用的PyTorch版本低于1.10,没有上述两个参数,只能沿用原有的通过onnx库二次读写的方案,建议优先升级PyTorch版本以使用更简洁的实现。
内容的提问来源于stack exchange,提问作者this_josh
相关产品推荐
相关产品推荐

