TensorFlow SavedModel转GraphDef后导入新Session报错,如何解决?
你遇到的问题核心在于GraphDef只保存了计算图的结构,并没有包含SavedModel里训练好的变量值。当你从原Session导出graph_def = sess.graph.as_graph_def()时,这个GraphDef里只有变量的定义(比如dense_1/kernel),但没有存储它的实际权重值——这些值是存在SavedModel的checkpoint文件里的。
当你用tf.import_graph_def导入这个图结构后,调用sess.run(tf.global_variables_initializer())只是给变量赋了默认的初始化值(比如随机值),但这并不是你原来SavedModel里训练好的权重,而且即使这样,有些变量可能在SavedModel里是通过特殊方式初始化的,单纯的全局初始化也没法正确加载原来的值,所以才会报错。
下面给你两种可行的解决方案:
方案一:直接在Spark中加载SavedModel(推荐)
既然你原本的模型是SavedModel格式,其实不需要转成GraphDef再导入,Spark可以直接加载SavedModel进行分布式推理。你可以在Spark的UDF里直接加载SavedModel:
def predict(image): with tf.Session() as sess: tf.saved_model.loader.load(sess, ['serve'], folder_path) logits = sess.run('dense_1/Softmax:0', {'input_1:0': [image]}) return logits[0] # 注册为Spark UDF predict_udf = udf(predict, ArrayType(FloatType())) df = df.withColumn('prediction', predict_udf('input_image_col'))
这种方式直接复用SavedModel的完整内容,不需要处理变量初始化的问题,而且能保证权重是训练好的原值。
方案二:导出GraphDef时同时保存变量值,导入后加载
如果你一定要用GraphDef的方式,需要把原Session里的变量值一起保存,然后在新Session里导入GraphDef后再加载这些值:
步骤1:保存GraphDef和变量值
在原Session中,除了导出GraphDef,还要把变量值保存成checkpoint:
sess = tf.Session() tf.saved_model.loader.load(sess, ['serve'], folder) # 保存GraphDef tf.train.write_graph(sess.graph_def, './model_dir', 'model_graph.pb', as_text=False) # 保存变量值 saver = tf.train.Saver() saver.save(sess, './model_dir/model.ckpt')
步骤2:导入GraphDef并加载变量值
在新Session中,导入GraphDef后用Saver加载之前保存的checkpoint:
with tf.Session(graph=tf.Graph()) as sess: # 导入GraphDef with tf.gfile.GFile('./model_dir/model_graph.pb', 'rb') as f: graph_def = tf.GraphDef() graph_def.ParseFromString(f.read()) tf.import_graph_def(graph_def, name="") # 加载变量值 saver = tf.train.Saver() saver.restore(sess, './model_dir/model.ckpt') # 现在可以正常运行推理了 result = sess.run('dense_1/Softmax:0', {'input_1:0': input_image})
这样就能保证导入的变量是原来训练好的权重,不会出现未初始化的错误。
另外要注意:如果你的SavedModel里有一些资源型变量(比如LookupTable),这种方式可能还需要额外处理,所以优先推荐方案一直接加载SavedModel。
内容的提问来源于stack exchange,提问作者user3599803

