You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.26 10:04:38