TensorFlow加载训练模型时出现KeyError: 'BlockLSTM'问题求助
这个错误的核心原因是:你的TensorFlow环境无法识别BlockLSTM这个操作符(op)。BlockLSTM并非TensorFlow核心API中的原生op,它通常来自tensorflow.contrib.rnn模块(在TensorFlow 1.x版本中),或者是你自定义的LSTM变体实现。当你尝试通过import_meta_graph加载模型时,当前环境没有注册这个op的定义,就会触发KeyError。
下面是针对性的解决方案:
1. 确保TensorFlow版本完全一致
模型保存和加载必须使用完全相同的TensorFlow版本,尤其是TensorFlow 1.x系列中,contrib模块的内容在不同小版本间有较大变动。比如你保存模型用的是TF1.13,加载时就不能用TF1.15,否则可能出现op不兼容的情况。
2. 提前导入包含BlockLSTM的模块
在执行import_meta_graph之前,先导入定义BlockLSTM的模块,确保TensorFlow能识别这个op。如果你的BlockLSTM来自contrib.rnn,在加载代码开头添加:
import tensorflow as tf import tensorflow.contrib.rnn as rnn # 显式导入contrib.rnn模块,注册BlockLSTM op
3. 优先通过重建网络结构加载模型(更可靠)
直接使用import_meta_graph容易受op注册问题影响,更稳妥的方式是先重建和训练时完全一致的网络结构,再通过Saver加载权重:
# 第一步:先重建你的模型结构(和训练时的代码完全一致,包括BlockLSTM的定义) from your_model_definition_file import YourModel # 假设你的模型类在这个文件里 c = ... # 你的配置参数 model = YourModel(c) # 第二步:加载权重 tf.reset_default_graph() with tf.Session(graph=model.graph) as sess: saver = tf.train.Saver() saver.restore(sess, tf.train.latest_checkpoint('./')) # 后续执行推理 pred = sess.run(model.preds, feed_dict={model.input_data: model_input})
4. 检查是否是自定义BlockLSTM
如果BlockLSTM是你自己实现的自定义op(比如C++编译的.so文件),那么在加载模型前必须先加载这个自定义op的库:
tf.load_op_library('path/to/your_block_lstm_op.so')
如果是Python层面自定义的LSTM类,要确保在加载模型前已经执行了该类的定义代码。
结合你的代码,最快速的修复应该是在模型恢复代码开头添加import tensorflow.contrib.rnn as rnn,同时确保TF版本和训练时一致。
内容的提问来源于stack exchange,提问作者Logen Base

