【发布时间】:2016-10-05 13:36:54
【问题描述】:
我认为,如果针对 convnet in the CIFAR-10 tutorial 创建的模型测试单个新图像这一关键任务有一个有据可查的解决方案,这将对 Tensorflow 社区大有帮助。
我可能错了,但似乎缺少使训练模型在实践中可用的关键步骤。该教程中有一个“缺失的环节”——一个脚本可以直接加载单个图像(作为数组或二进制),将其与经过训练的模型进行比较,然后返回一个分类。
先前的答案给出了解释整体方法的部分解决方案,但我都无法成功实施。可以在这里和那里找到其他零碎的东西,但不幸的是还没有添加到一个有效的解决方案中。在将其标记为重复或已回答之前,请考虑我所做的研究。
Tensorflow: how to save/restore a model?
Unable to restore models in tensorflow v0.8
https://gist.github.com/nikitakit/6ef3b72be67b86cb7868
最流行的答案是第一个,其中@RyanSepassi 和@YaroslavBulatov 描述了问题和方法:需要“手动构建具有相同节点名称的图,并使用 Saver 将权重加载到其中”。尽管这两个答案都有帮助,但如何将其插入 CIFAR-10 项目尚不清楚。
非常需要一个功能齐全的解决方案,因此我们可以将其移植到其他单一图像分类问题。在这方面有几个关于 SO 的问题要求这个,但仍然没有完整的答案(例如Load checkpoint and evaluate single image with tensorflow DNN)。
我希望我们能集中在一个每个人都可以使用的工作脚本上。
以下脚本尚不可用,我很高兴收到您关于如何改进此脚本以使用 CIFAR-10 TF 教程训练模型提供单图像分类解决方案的意见。
假设所有变量、文件名等都未触及原始教程。
新文件:cifar10_eval_single.py
import cv2
import tensorflow as tf
FLAGS = tf.app.flags.FLAGS
tf.app.flags.DEFINE_string('eval_dir', './input/eval',
"""Directory where to write event logs.""")
tf.app.flags.DEFINE_string('checkpoint_dir', './input/train',
"""Directory where to read model checkpoints.""")
def get_single_img():
file_path = './input/data/single/test_image.tif'
pixels = cv2.imread(file_path, 0)
return pixels
def eval_single_img():
# below code adapted from @RyanSepassi, however not functional
# among other errors, saver throws an error that there are no
# variables to save
with tf.Graph().as_default():
# Get image.
image = get_single_img()
# Build a Graph.
# TODO
# Create dummy variables.
x = tf.placeholder(tf.float32)
w = tf.Variable(tf.zeros([1, 1], dtype=tf.float32))
b = tf.Variable(tf.ones([1, 1], dtype=tf.float32))
y_hat = tf.add(b, tf.matmul(x, w))
saver = tf.train.Saver()
with tf.Session() as sess:
sess.run(tf.initialize_all_variables())
ckpt = tf.train.get_checkpoint_state(FLAGS.checkpoint_dir)
if ckpt and ckpt.model_checkpoint_path:
saver.restore(sess, ckpt.model_checkpoint_path)
print('Checkpoint found')
else:
print('No checkpoint found')
# Run the model to get predictions
predictions = sess.run(y_hat, feed_dict={x: image})
print(predictions)
def main(argv=None):
if tf.gfile.Exists(FLAGS.eval_dir):
tf.gfile.DeleteRecursively(FLAGS.eval_dir)
tf.gfile.MakeDirs(FLAGS.eval_dir)
eval_single_img()
if __name__ == '__main__':
tf.app.run()
【问题讨论】:
标签: python python-3.x machine-learning tensorflow