如何从tf.estimator训练的CNN获取预测?解决调用报错问题
问题分析与解决方案
首先直接点出你遇到的核心问题:你在Controller_tf类里调用model(self.x, None, tf.estimator.ModeKeys.PREDICT)后,把返回值当成张量去执行sess.run(),但你的cnn_model_fn3在PREDICT模式下返回的是tf.estimator.EstimatorSpec对象(而不是张量),这就导致TensorFlow无法识别这个对象为可获取的计算节点,从而抛出了Fetch argument 'infer' cannot be interpreted as a Tensor的错误。另外你的模型函数末尾还有一段多余的if mode == PREDICT返回logits的代码,这段代码永远不会被执行(因为前面已经返回了EstimatorSpec),需要删掉。
下面给你两种可行的解决方案,你可以根据自己的需求选择:
方案一:使用tf.estimator原生的predict接口(推荐)
tf.estimator已经封装了Session管理、模型加载等逻辑,不需要手动创建Session和Saver,用它的predict方法更符合框架设计规范,也更简洁。
步骤1:修正模型函数
先把模型函数里多余的代码删掉,确保PREDICT分支正确返回包含所需张量的EstimatorSpec:
def cnn_model_fn3(features, labels, mode): # 统一处理输入层 if mode == tf.estimator.ModeKeys.PREDICT: input_layer = features # 预测时features直接是输入张量 else: input_layer = tf.reshape(features["image_data"], [-1, 104, 160, 3]) conv1 = tf.layers.conv2d( inputs=input_layer, filters=32, kernel_size=[10, 10], padding="same", activation=tf.nn.relu, name='Conv1') # ... 保留你其他的层代码 ... logits = tf.layers.dense( inputs=dropout1, units=3, name='Dense3') predictions = { "classes": tf.argmax(input=logits, axis=1), "probabilities": tf.nn.softmax(logits, name="softmax_tensor"), "logits": logits # 把logits加入预测结果,方便后续获取 } if mode == tf.estimator.ModeKeys.PREDICT: return tf.estimator.EstimatorSpec(mode=mode, predictions=predictions) # 训练和评估模式的逻辑保持不变 loss = tf.losses.sparse_softmax_cross_entropy(labels=labels, logits=logits) if mode == tf.estimator.ModeKeys.TRAIN: optimizer = tf.train.GradientDescentOptimizer(learning_rate=0.001) train_op = optimizer.minimize( loss=loss, global_step=tf.train.get_global_step()) return tf.estimator.EstimatorSpec(mode=mode, loss=loss, train_op=train_op) eval_metric_ops = { "accuracy": tf.metrics.accuracy( labels=labels, predictions=predictions["classes"])} return tf.estimator.EstimatorSpec( mode=mode, loss=loss, eval_metric_ops=eval_metric_ops)
步骤2:修改Controller_tf类适配estimator
import os import tensorflow as tf import numpy as np class Controller_tf: set_speed = None def __init__(self, model_fn, ckpt_dir, set_speed_in): self.set_speed = set_speed_in # 创建Estimator对象,指定模型函数和模型目录(会自动加载最新检查点) self.estimator = tf.estimator.Estimator( model_fn=model_fn, model_dir=ckpt_dir ) def update(self, message): # 处理输入图像 image = frame2numpy(message['frame'], (160,104)) image_array = np.asarray(image) # 构建输入数据,匹配模型PREDICT模式的输入格式 input_fn = tf.estimator.inputs.numpy_input_fn( x=image_array[None, :, :, :], # 增加batch维度 shuffle=False ) # 获取预测结果(取第一个结果,因为batch size是1) predictions = next(self.estimator.predict(input_fn=input_fn)) # 返回你需要的结果,这里返回logits,也可以返回classes或probabilities return predictions["logits"]
步骤3:修改调用代码
model = cnn_model_fn3 ckpt_dir = 'ckpts/stc_model3/' # 这里传模型目录,不是具体的ckpt文件 controller = Controller_tf(model, ckpt_dir, 18) image_file = 'G:/Datasets/ds072.001/ds072.001-fm-0008465.jpg' satnavimg = load_image(image_file) satnavimg = (satnavimg/127.5) - 1.0 msg = {'frame': satnavimg} turn = controller.update(msg) print(turn)
方案二:手动用Session加载模型(兼容旧代码结构)
如果你想保留原来的Session调用逻辑,可以把模型的网络结构抽离成独立函数,避免直接调用返回EstimatorSpec的模型函数。
步骤1:抽离预测网络结构
def build_prediction_network(input_tensor): input_layer = input_tensor conv1 = tf.layers.conv2d( inputs=input_layer, filters=32, kernel_size=[10, 10], padding="same", activation=tf.nn.relu, name='Conv1') # ... 完全复制你模型函数里的所有层代码 ... logits = tf.layers.dense( inputs=dropout1, units=3, name='Dense3') return logits
步骤2:修改模型函数调用抽离的网络
def cnn_model_fn3(features, labels, mode): if mode == tf.estimator.ModeKeys.PREDICT: input_layer = features else: input_layer = tf.reshape(features["image_data"], [-1, 104, 160, 3]) # 调用抽离的预测网络 logits = build_prediction_network(input_layer) predictions = { "classes": tf.argmax(input=logits, axis=1), "probabilities": tf.nn.softmax(logits, name="softmax_tensor"), "logits": logits } # 后续的EstimatorSpec返回逻辑和方案一的模型函数一致...
步骤3:修改Controller_tf类
class Controller_tf: set_speed = None def __init__(self, ckpt_path, set_speed_in): self.set_speed = set_speed_in self.x = tf.placeholder(tf.float32, shape=(None, 104, 160, 3)) self.y = build_prediction_network(self.x) # 调用抽离的网络,得到张量 # 配置Session config = tf.ConfigProto() config.gpu_options.allow_growth = True self.sess = tf.Session(config=config) # 加载检查点 saver = tf.train.Saver() saver.restore(self.sess, ckpt_path) def update(self, message): image = frame2numpy(message['frame'], (160,104)) image_array = np.asarray(image) turn_logits = self.sess.run(self.y, {self.x: image_array[None, :, :, :]}) return turn_logits
这样修改后,你原来的调用代码基本不需要改动,就能正常获取预测结果了。
内容的提问来源于stack exchange,提问作者tinyMind
相关产品推荐
相关产品推荐

