加载MobileFaceNet预训练模型报错:Cannot parse file 'saved_model.pb'排查
加载MobileFaceNet预训练模型为Keras对象的问题解决
问题根源
你碰到的这个报错,核心原因是这个GitHub仓库里的预训练模型不是Keras原生导出的格式。tf.keras.models.load_model()只能加载用model.save()或者tf.keras.models.save_model()保存的Keras模型(标准SavedModel或者.h5格式),但这个仓库里的模型是用TensorFlow原生静态图API构建并保存的(可能是冻结图或者checkpoint格式),两者的文件结构和内部格式不兼容,所以无法直接解析。
正确加载并使用模型的两种方法
方法一:重构Keras模型结构 + 加载预训练权重(推荐,完全支持.fit()和.predict())
这种方法能得到标准的Keras模型,后续不管是推理还是继续训练都没问题:
- 复刻模型结构:去仓库里找到MobileFaceNet的TensorFlow定义代码,把它转换成tf.keras的实现。比如用
tf.keras.layers.Conv2D替代原生TF的tf.nn.conv2d,用tf.keras.layers.BatchNormalization替换TF的批量归一化层,确保层的参数(比如卷积核大小、步长、通道数)和原模型完全一致。 - 加载预训练权重:
如果模型是checkpoint格式(目录里有.index、.data-00000-of-00001和checkpoint文件),可以用TensorFlow的Checkpoint工具加载:import tensorflow as tf # 假设你已经用Keras定义好了模型mobile_face_net checkpoint = tf.train.Checkpoint(model=mobile_face_net) # 加载最新的checkpoint checkpoint.restore(tf.train.latest_checkpoint("你的模型目录路径")).expect_partial() # expect_partial()用来忽略一些不影响推理/训练的变量(比如原模型里的优化器状态) - 验证模型:随便找一张测试图输入
model.predict(),看输出是否符合预期,没问题的话就可以正常用.fit()继续训练了。
方法二:加载静态图并封装为Keras模型(适合快速推理,可能无法直接.fit())
如果不想重构模型,可以把原模型的静态图封装成Keras模型,适合快速做推理:
import tensorflow as tf def load_mobilefacenet_as_keras_model(frozen_pb_path): # 加载冻结的.pb图文件 with tf.io.gfile.GFile(frozen_pb_path, 'rb') as f: graph_def = tf.compat.v1.GraphDef() graph_def.ParseFromString(f.read()) # 把图导入到当前计算图 with tf.Graph().as_default() as graph: tf.import_graph_def(graph_def, name='') # 获取模型的输入和输出张量(需要替换成原模型实际的节点名) # 你可以用tf.compat.v1.get_default_graph().get_operations()查看所有节点 input_tensor = graph.get_tensor_by_name('input:0') output_tensor = graph.get_tensor_by_name('embeddings:0') # 封装成Keras模型 keras_model = tf.keras.Model(inputs=input_tensor, outputs=output_tensor) return keras_model # 调用示例 model = load_mobilefacenet_as_keras_model("path/to/frozen_inference_graph.pb") # 现在可以用model.predict()做推理
注意:这种方法得到的模型是基于静态图封装的,没有保存训练相关的计算节点(比如损失函数、优化器),所以如果需要继续训练,还是建议用方法一。
内容的提问来源于stack exchange,提问作者Trọng Nghĩa Lê Đình
相关产品推荐
相关产品推荐

