将Tensor保存为TensorProto时无SerializeToString属性报错如何解决
问题原因
- 你当前使用的变量
x是TensorFlow框架的Tensor类型,而非ONNX定义的TensorProto类型,SerializeToString是ONNXTensorProto类独有的序列化方法,TensorFlow的Tensor对象没有实现该方法,因此触发属性不存在的报错。
解决方法
按照以下步骤转换类型后再保存即可:
- 先将TensorFlow的Tensor转换为numpy数组
# Eager Execution(TensorFlow 2.x默认模式)下直接调用numpy方法 x_np = x.numpy() # 若为TensorFlow 1.x的Graph模式,需要先在会话中求值再转numpy # with tf.Session() as sess: # x_np = sess.run(x)
- 借助ONNX提供的工具方法将numpy数组转为
TensorProto对象
import onnx from onnx import numpy_helper # 转换为TensorProto,可自定义name参数设置张量名 x_tensor_proto = numpy_helper.from_array(x_np, name="images")
- 执行序列化保存操作
with open('tensor.pb', 'wb+') as f: f.write(x_tensor_proto.SerializeToString())
注意:若运行时提示缺少onnx依赖,可先执行
pip install onnx完成安装。
内容的提问来源于stack exchange,提问作者harry
相关产品推荐
相关产品推荐

