如何在Python中为DeepStream实现NHWC到NCHW格式转换?
解决方案:TensorFlow PB转ONNX时将输入格式从NHWC转为NCHW
针对tf2onnx转换后输入仍为NHWC、手动转置代码未生效的问题,可按以下步骤解决:
步骤1:确认输入节点的准确名称
你之前使用的input0:0可能不是模型实际的输入节点名,导致--inputs-as-nchw参数失效。执行以下命令查看SavedModel的输入节点信息:
saved_model_cli show --dir model.pb --all
在输出的MetaGraphDef with tag-set: 'serve' contains the following SignatureDefs:部分,找到serving_default下的输入节点名称(例如可能是input_1,而非input0:0)。
步骤2:升级tf2onnx并重新转换
旧版本tf2onnx可能存在--inputs-as-nchw参数兼容问题,先升级工具:
pip install --upgrade tf2onnx
然后使用正确的输入节点名执行转换命令:
!python -m tf2onnx.convert --saved-model model.pb --output model_nchw.onnx --inputs-as-nchw 你的输入节点名
(替换你的输入节点名为步骤1中获取的名称,无需加:0后缀)
步骤3:修改TensorFlow模型后重新导出(如果步骤2仍失效)
如果上述方法无效,可手动修改模型结构,将输入层改为NCHW格式后再转ONNX:
import tensorflow as tf # 加载原SavedModel模型 original_model = tf.keras.models.load_model('model.pb') # 定义NCHW格式的输入(shape为[batch, channels, height, width]) nchw_input = tf.keras.Input(shape=(3, 200, 300), name='input_nchw') # 将NCHW转置为原模型需要的NHWC格式 nhwc_input = tf.transpose(nchw_input, perm=[0, 2, 3, 1]) # 连接原模型 model_outputs = original_model(nhwc_input) # 创建新的模型,输入为NCHW格式 new_model = tf.keras.Model(inputs=nchw_input, outputs=model_outputs) # 保存修改后的SavedModel new_model.save('model_nchw_input')
再用tf2onnx转换新模型:
!python -m tf2onnx.convert --saved-model model_nchw_input --output model_nchw.onnx
验证结果
使用Netron工具打开转换后的model_nchw.onnx,检查输入节点的shape是否为[batch_size, 3, 200, 300](NCHW格式)。
内容的提问来源于stack exchange,提问作者Sovik Gupta
相关产品推荐
相关产品推荐

