TensorFlow中CNN权重偏置重加载报错及可视化问题咨询
解决TensorFlow加载模型时的InvalidArgumentError及计算图可视化方法
一、解决占位符未喂入数据的错误
你遇到的这个报错逻辑其实很清晰:y_pred这个输出节点的计算依赖占位符x的输入,但你加载模型后直接调用sess.run('y_pred:0'),完全没给x传数据,TensorFlow自然会提示你必须喂入符合要求的张量值。
解决起来分两步走:
获取图中的占位符和输出节点
导入meta图后,所有节点都存放在默认计算图里,你可以通过节点名称把它们取出来:import tensorflow as tf import numpy as np sess = tf.Session() saver = tf.train.import_meta_graph('results/steering_model.meta') saver.restore(sess, 'results/steering_model') # 获取默认图对象 graph = tf.get_default_graph() # 通过名称获取占位符x(注意要加:0,这是TensorFlow节点张量的命名规则) x = graph.get_tensor_by_name('x:0') # 获取输出节点y_pred y_pred = graph.get_tensor_by_name('y_pred:0')准备符合形状的输入数据并喂入
你的x形状是[16,96,128,3],意味着需要传入16张96×128的3通道彩色图(float32类型)。可以用随机数快速测试,或者加载预处理后的真实测试图片:# 生成符合要求的测试输入(随机数仅用于测试,实际用真实图片需和训练时做同样预处理) test_input = np.random.rand(16, 96, 128, 3).astype(np.float32) # 喂入数据并运行得到预测结果 pred_result = sess.run(y_pred, feed_dict={x: test_input}) print(pred_result)要是用真实图片,记得要做和训练阶段完全一致的预处理(比如归一化到0-1、尺寸裁剪等),否则输入格式不匹配也会报错。
二、计算图可视化方法
用TensorFlow自带的TensorBoard就能轻松可视化计算图,不管是训练时的图还是加载后的图都能搞定:
方法1:训练时保存图到日志
在你的训练代码里,创建Session后添加这两行,就能把计算图写入日志文件:
# 训练代码片段中添加 writer = tf.summary.FileWriter('./training_graph_logs', session.graph) writer.close()
方法2:加载模型后保存图到日志
如果是已经训练好的模型,加载后同样可以写入日志:
# 加载模型的代码里添加 writer = tf.summary.FileWriter('./loaded_graph_logs', sess.graph) writer.close()
启动TensorBoard查看
打开终端,运行以下命令(替换成你的日志目录):
tensorboard --logdir=./training_graph_logs
然后打开浏览器,访问终端输出的地址(通常是http://localhost:6006),点击左侧的Graphs标签,就能看到完整的计算图结构了——可以缩放、展开节点,查看每一层的细节和依赖关系。
内容的提问来源于stack exchange,提问作者Federico Màlato
相关产品推荐
相关产品推荐

