无法使用训练后的TensorFlow模型,求解决方案及规避ParseFromString方法
解决InceptionV3转saved_model.pb调用错误及规避
graph_def.ParseFromString()的方法 嘿,作为刚接触深度学习和TensorFlow的新手,遇到这种模型加载报错的问题肯定超崩溃吧!我刚入门的时候也踩过类似的坑,给你分享几个靠谱的解决方法,还能帮你避开那个麻烦的graph_def.ParseFromString()函数~
一、先确认你的saved_model.pb生成是否合规
很多时候调用出错的根源是模型导出环节出了问题,先检查下你是不是用了正确的方式导出SavedModel:
- 别手动去处理GraphDef!用TensorFlow官方推荐的
tf.saved_model.save()导出,代码示例如下:
这样导出的模型是完整的TensorFlow标准格式,后续加载完全不需要碰import tensorflow as tf # 假设你已经完成InceptionV3的微调,加载了自己训练的权重 model = tf.keras.applications.InceptionV3( weights='your_finetuned_weights.h5', include_top=True, classes=你的类别数量 ) # 正确导出SavedModel(会生成包含saved_model.pb的目录) tf.saved_model.save(model, "./my_saved_model")graph_def.ParseFromString()。
二、规避graph_def.ParseFromString()的加载方法
如果你已经有了saved_model.pb,直接用下面两种简单的方式加载就行,完全绕开那个容易出错的函数:
方法1:用Keras一键加载(最适合新手)
Keras的API已经封装了所有复杂操作,一行代码就能加载模型并直接推理:
import tensorflow as tf from tensorflow.keras.preprocessing import image # 加载SavedModel为Keras模型 loaded_model = tf.keras.models.load_model("./my_saved_model") # 预处理测试图像(InceptionV3要求输入尺寸是299x299) img = image.load_img("test_image.jpg", target_size=(299, 299)) img_array = image.img_to_array(img) img_array = tf.expand_dims(img_array, 0) # 增加batch维度 # 直接预测 predictions = loaded_model.predict(img_array)
方法2:用SavedModelLoader加载(适配低版本TensorFlow)
如果你的TensorFlow版本较低(比如1.x),可以用这个底层一点的方法,同样不需要解析GraphDef:
import tensorflow as tf with tf.Session() as sess: # 加载SavedModel,指定SERVING标签 tf.saved_model.loader.load( sess, [tf.saved_model.tag_constants.SERVING], "./my_saved_model" ) # 获取模型的输入输出张量(InceptionV3默认输入名是input_1:0,输出名是predictions/Softmax:0) input_tensor = sess.graph.get_tensor_by_name("input_1:0") output_tensor = sess.graph.get_tensor_by_name("predictions/Softmax:0") # 预处理图像后传入推理 img = ... # 把你的图像转换成符合要求的张量格式 predictions = sess.run(output_tensor, feed_dict={input_tensor: img})
三、如果一定要用graph_def.ParseFromString()(排查错误)
要是你因为某些特殊需求必须用这个函数,那先排查这两个常见错误:
- 文件读取模式错误:必须用二进制模式读取saved_model.pb,不能用文本模式:
# 错误写法(文本模式读取二进制文件会损坏内容) with open("saved_model.pb", "r") as f: graph_def.ParseFromString(f.read()) # 正确写法 with open("saved_model.pb", "rb") as f: graph_def.ParseFromString(f.read()) - 混淆了SavedModel和GraphDef格式:saved_model.pb里包含的是MetaGraphDef,不是单纯的GraphDef,得先解析MetaGraph再提取GraphDef:
from tensorflow.core.protobuf import saved_model_pb2 saved_model = saved_model_pb2.SavedModel() with open("saved_model.pb", "rb") as f: saved_model.ParseFromString(f.read()) # 从第一个MetaGraph中获取GraphDef graph_def = saved_model.meta_graphs[0].graph_def
总结
作为新手,真心建议你直接用前面两种规避graph_def.ParseFromString()的方法,官方封装好的API已经帮你处理了所有底层的图解析和变量恢复工作,既简单又不容易出错,能省掉超多调试时间~
内容的提问来源于stack exchange,提问作者Torab Shaikh
相关产品推荐
相关产品推荐

