TensorFlow MNIST模型保存恢复报错:placeholder_1需输入值
问题分析与修复方案
我来帮你梳理下问题根源,以及对应的解决办法:
错误原因
你在恢复模型时犯了两个核心错误:
- 模型结构不匹配:你重新定义了
x和y作为占位符,但原模型中的y是通过tf.nn.softmax(tf.matmul(x,W) + b)计算出的输出张量,并非占位符;同时原模型里用于训练的真实标签占位符y_(也就是报错里的Placeholder_1)依然存在于计算图中,当你执行sess.run(y)时,TensorFlow会检查整个计算图的依赖,发现y_未被喂入数据,因此抛出错误。 - Meta图保存缺失:你训练时使用了
write_meta_graph=False参数,导致没有生成model_mnist.ckpt.meta文件,后续尝试用tf.train.import_meta_graph加载会直接失败。
修复方案
方案一:重新定义与训练一致的模型结构(无需Meta图)
这种方法不需要依赖Meta图文件,只需完全复刻训练时的模型结构(跳过训练相关的y_和优化器操作),再加载保存的变量即可:
import cv2 import numpy as np import tensorflow as tf # 图片预处理(和训练时逻辑一致) img = cv2.imread('lena.png') img = cv2.resize(img, (28,28)) img = cv2.cvtColor(img,cv2.COLOR_RGB2GRAY) arr = [] for i in range(28): for j in range(28): gray = 1 - img[i,j]/255 arr.append(gray) arr_mnist = np.array([arr]) # 复刻训练时的模型结构(仅保留推理必要部分) tf.reset_default_graph() x = tf.placeholder(tf.float32, [None, 784]) W = tf.Variable(tf.zeros([784,10])) b = tf.Variable(tf.zeros([10])) y = tf.nn.softmax(tf.matmul(x,W) + b) # 恢复模型变量 sess = tf.Session() saver = tf.train.Saver() saver.restore(sess, './model_mnist.ckpt') # 执行推理(仅需喂入输入x的数据) result = sess.run(y, feed_dict={x: arr_mnist}) print(result) print("预测值为:",np.argmax(result[0]),";概率为:",np.max(result[0])/np.sum(result[0]))
方案二:修改训练代码保存Meta图,直接加载完整计算图
如果你想直接加载训练时的计算图,需要先修改训练阶段的保存代码,让它生成Meta图文件:
# 训练时的保存代码修改为(去掉write_meta_graph=False) save_path = saver.save(sess, "./model_mnist.ckpt")
之后使用以下代码恢复模型:
import cv2 import numpy as np import tensorflow as tf # 图片预处理 img = cv2.imread('lena.png') img = cv2.resize(img, (28,28)) img = cv2.cvtColor(img,cv2.COLOR_RGB2GRAY) arr = [] for i in range(28): for j in range(28): gray = 1 - img[i,j]/255 arr.append(gray) arr_mnist = np.array([arr]) # 加载完整计算图与变量 tf.reset_default_graph() sess = tf.Session() saver = tf.train.import_meta_graph('./model_mnist.ckpt.meta') saver.restore(sess, './model_mnist.ckpt') # 从图中获取输入x和输出y的张量 graph = tf.get_default_graph() x = graph.get_tensor_by_name("Placeholder:0") # 原模型中x的默认名称 y = graph.get_tensor_by_name("Softmax:0") # 原模型中softmax输出的默认名称 # 执行推理 result = sess.run(y, feed_dict={x: arr_mnist}) print(result) print("预测值为:",np.argmax(result[0]),";概率为:",np.max(result[0])/np.sum(result[0]))
注:如果不确定张量名称,可以在训练时通过
print(x.name)和print(y.name)查看具体名称,再替换到恢复代码中。
内容的提问来源于stack exchange,提问作者xuchenxing
相关产品推荐
相关产品推荐

