TensorFlow测试集预测输出问题求助:sess.run输入参数填写
解决TensorFlow中获取测试数据预测结果的sess.run参数问题
看起来你在训练完模型后,想获取测试数据的预测结果,但卡在了sess.run()的参数上。我来帮你梳理一下正确的步骤:
首先,先修正一个代码里的小问题——你的nn_model函数目前没有返回值,这样logits = nn_model(next_element)会得到None,后续的tf.argmax肯定会报错。你需要在函数末尾返回最终的网络输出(也就是logits),比如:
def nn_model(data): convnet = conv_2d(data, 32, 3, padding='same', activation='relu') # 注意这里参数应该是传入的data,不是in_data convnet = max_pool_2d(convnet, 2) # 如果后面还有全连接层或者其他输出层,记得在这里返回最终的logits return convnet
接下来是获取预测结果的正确步骤:
- 首先,你需要初始化测试数据的迭代器。从代码里的
training_init_op来看,你应该是用了数据集迭代器加载训练数据,那测试数据肯定也有对应的初始化操作(比如test_init_op,你需要自己定义这个,对应测试数据集的迭代器初始化)。在获取预测前,先运行这个操作:
sess.run(test_init_op)
- 然后,直接把你定义的
prediction作为参数传入sess.run()就可以获取预测结果了。如果是批量处理测试数据,可以用循环来读取所有结果:
# 方式1:获取单批次测试数据的预测结果 batch_pred = sess.run(prediction) print("当前批次预测结果:", batch_pred) # 方式2:获取所有测试数据的预测结果 all_predictions = [] try: while True: pred = sess.run(prediction) all_predictions.extend(pred) except tf.errors.OutOfRangeError: # 迭代器到末尾时会抛出这个异常,代表所有测试数据都处理完了 print("所有测试数据预测完成") print("全部预测结果:", all_predictions)
为什么这样可行?因为你的prediction是基于logits,而logits又依赖于next_element(测试数据迭代器输出的样本),当你运行sess.run(prediction)时,TensorFlow会自动计算所有依赖的张量——从迭代器中获取测试数据,经过网络得到logits,最后输出argmax后的预测类别。
最后提醒一下:训练完后一定要切换到测试迭代器,不然会继续从训练数据里读取样本,得到的就不是测试数据的预测结果了。
内容的提问来源于stack exchange,提问作者J Wu
相关产品推荐
相关产品推荐

