- html - 出于某种原因,IE8 对我的 Sass 文件中继承的 html5 CSS 不友好?
- JMeter 在响应断言中使用 span 标签的问题
- html - 在 :hover and :active? 上具有不同效果的 CSS 动画
- html - 相对于居中的 html 内容固定的 CSS 重复背景?
我正在使用 Tensorflow 开发一个项目。我已经构建并训练了一个 CNN,现在我尝试将其加载到另一个文件中以进行预测。由于某种原因,我不断收到错误“您必须使用 dtype float 和 shape [10] 为占位符张量 'y_pred' 提供值”
构建图表的文件有一个用于预测的变量 y_pred:
y_pred = tf.nn.softmax(layer_fc2)
我尝试加载模型的文件如下:
# Create Session
sess = tf.Session()
# Load model
saver = tf.train.import_meta_graph('Model.meta')
saver.restore(sess, tf.train.latest_checkpoint('./'))
sess.run(tf.global_variables_initializer())
graph = tf.get_default_graph()
x_batch = mnist.test.next_batch(1)
x_batch = x_batch[0].reshape(1, 784)
x = graph.get_tensor_by_name("x:0")
y_pred = graph.get_tensor_by_name("y_pred:0")
classification = sess.run(y_pred, feed_dict={x:x_batch})
print(classification)
我收到的确切错误是:
InvalidArgumentError (see above for traceback): You must feed a value for placeholder tensor 'y_pred' with dtype float and shape [10]
[[Node: y_pred = Placeholder[dtype=DT_FLOAT, shape=[10], _device="/job:localhost/replica:0/task:0/cpu:0"]()]]
我想知道在导出之前是否我没有正确设置该值?有谁知道为什么这不起作用?
编辑。包括型号代码:
# Network Design
# First Layer
layer_conv1, weights_conv1 = new_conv_layer(input=x_image, num_input_channels=num_channels, filter_size=filter_size1, num_filters=num_filters1, use_pooling=True)
# Second Layer
layer_conv2, weights_conv2 = new_conv_layer(input=layer_conv1, num_input_channels=num_filters1, filter_size=filter_size2, num_filters=num_filters2, use_pooling=True)
# Third Layer
layer_conv3, weights_conv3 = new_conv_layer(input=layer_conv2, num_input_channels=num_filters2, filter_size=filter_size3, num_filters=num_filters3, use_pooling=True)
# Flatten Layer
layer_flat, num_features = flatten_layer(layer_conv3)
# First Fully Connected Layer
layer_fc1 = new_fc_layer(input=layer_flat, num_inputs=num_features, num_outputs=fc_size, use_relu=True)
# Second Fully Connected Layer
layer_fc2 = new_fc_layer(input=layer_fc1, num_inputs=fc_size, num_outputs=num_classes, use_relu=False)
# softmaxResult = tf.placeholder(tf.float32, shape=[10], name='softmaxResult')
# Get class probabilities
y_pred = tf.nn.softmax(layer_fc2)
y_pred = tf.identity(y_pred, name="y_pred")
# session.run(y_pred, feed_dict={softmaxResult: y_pred})
# Predicted Class
y_pred_cls = tf.argmax(y_pred, dimension=1)
# softmaxResult.assign(y_pred_cls)
# Feed y_pred
# session.run(softmaxResult, feedDict={softmaxResult: softmaxResult})
# Define Cost Function
cross_entropy = tf.nn.softmax_cross_entropy_with_logits(logits=layer_fc2, labels=y_true)
cost = tf.reduce_mean(cross_entropy)
# Optimize Network
optimizer = tf.train.AdamOptimizer(learning_rate=1e-4).minimize(cost)
correct_prediction = tf.equal(y_pred_cls, y_true_cls)
accuracy = tf.reduce_mean(tf.cast(correct_prediction, tf.float32))
# Run Session
session.run(tf.global_variables_initializer())
def print_progress(epoch, feed_dict_train, feed_dict_validate, val_loss):
#Calculate accuracy on training set
acc = session.run(accuracy, feed_dict=feed_dict_train)
val_acc = session.run(accuracy, feed_dict=feed_dict_validate)
msg = "Epoch {0} --- Training Accuracy: {1:>6.1%}, Validation Accuracy: {2:>6.1%}, Validation Loss: {3:.3f}"
print(msg.format(epoch + 1, acc, val_acc, val_loss))
total_iterations = 0
#Optimization Function
def optimize(num_iterations):
# Updates global rather than local value
global total_iterations
best_val_loss = float("inf")
for i in range(total_iterations, total_iterations + num_iterations):
# Get training data batch
x_batch, y_batch = mnist.train.next_batch(batch_size)
# Get a validation batch
x_validate, y_validate = mnist.train.next_batch(batch_size)
# Shrink to single dimension
x_batch = x_batch.reshape(batch_size, img_size_flat)
x_validate = x_validate.reshape(batch_size, img_size_flat)
# Training feed
feed_dict_train = {x: x_batch, y_true: y_batch}
feed_dict_validate = {x: x_validate, y_true: y_validate}
# Run the optimizer
session.run(optimizer, feed_dict=feed_dict_train)
# Print status at end of each epoch (defined as full pass through training dataset).
if i % int(5000/batch_size) == 0:
val_loss = session.run(cost, feed_dict=feed_dict_validate)
epoch = int(i / int(5000/batch_size))
print_progress(epoch, feed_dict_train, feed_dict_validate, val_loss)
total_iterations += num_iterations
optimize(num_iterations=3000)
# Save the final model
saver = tf.train.Saver()
saved_path = saver.save(session, os.path.join(os.getcwd(),'MNIST Model'))
print("Model saved in: ", saved_path)
# Run on test image
image = mnist.test.next_batch(1)
feedin = image[0].reshape(1, 784)
inputStuff = {x:feedin}
classification = session.run(y_pred, feed_dict=inputStuff)
print(classification)
最佳答案
谢谢@VS_FF
您需要在“x:0”中找到他们输入的 key 。
关于python - Tensorflow 导入元图占位符未提供,我们在Stack Overflow上找到一个类似的问题: https://stackoverflow.com/questions/44375600/
当我这样做时... import numpy as np ...我可以使用它但是... import pprint as pp ...不能,因为我需要这样做... from pprint import
我第一次尝试将 OpenCV 用于 Python 3。要安装,我只需在终端中输入“pip3 install opencv-python”。当我这样做时,我在 Finder(我在 Mac 上)中看到,在
如果有一个库我将使用至少两种方法,那么以下之间在性能或内存使用方面是否有任何差异? from X import method1, method2 和 import X 最佳答案 有区别,因为在 imp
我正在从 lodash 导入一些函数,我的同事告诉我,单独导入每个函数比将它们作为一个组导入更好。 当前方法: import {fn1, fn2, fn3} from 'lodash'; 首选方法:
之间有什么关系: import WSDL 中的元素 -和- import元素和在 XML Schema ...尤其是 location 之间的关系前者和 schemaLocation 的属性后者的属性
我在从 'theano.configdefaults' 导入 'local_bitwidth' 时遇到问题。并显示以下消息: ImportError
我注意到 React 可以这样导入: import * as React from 'react'; ...或者像这样: import React from 'react'; 第一个导入 react
对于当前的项目,我必须使用矩阵中提供的信息并对其进行数学计算,以及使用 ITK/VTK 函数来显示医疗信息/渲染。基本上我必须以(我猜)50/50 的方式同时使用 matlab 例程和 VTK/ITK
当我看到 pysqlite 的示例时,SQLite 库有两个用例。 from sqlite3 import dbapi2 as sqlite3 和 import sqlite3 为什么有两种方式支持s
我使用 Anaconda Python 发行版:Python 2.7 x64 和 Windows 7 SP1 x64 Ultimate。 当我import matplotlib.pyplot时,我得到
目录 【容器】镜像导出/导入 导出 导入 带标签 不带标签,后期修改 【仓库】镜像导出/导入
我正在寻找一种导入模块的方法,以便我可以从子文件夹 project/v0 和根文件夹 project 运行脚本。/p> 我在 python 3.6 中的文件结构(这就是没有初始化文件的原因) proj
我通常被告知以下是不好的做法。 from module import * 主要原因(或者有人告诉我)是,您可能会导入一些您不想要的东西,并且它可能会隐藏另一个模块中具有类似名称的函数或类。 但是,Py
我为 urllib (python3) 编写了一个小包装器。在if中导入模块是否正确且安全? if self.response_encoding == 'gzip': import gzip
我正在 pimcore 中创建一个新站点。有没有办法导出/导入 pimcore 站点的完整数据,以便我可以导出 xml/csv 格式的 pimcore 数据进行必要的更改,然后将其导入回来? 最佳答案
在 Node JS 中测试以下模块布局,看起来本地导出的定义总是在名称冲突的情况下替换外部导出的定义(参见 B.js 中的 f1)。 A.js export const f1 = 'A' B.js e
我在使用 VBA 代码时遇到了一些问题,该代码应该将 excel 数据导入我的 Access 数据库。当我运行代码时,我收到一个运行时错误“运行时错误 438 对象不支持此属性或方法”。来自我在其他论
我有一个名为 elements 的包,其中包含按钮、trifader、海报等内容。在 Button 类中,我正在执行 from elements import * 这执行正常,当我尝试 print(p
在我长期使用 python 的经验中,我遇到了一个非常奇怪的问题。 提前我想说我想知道为什么会发生这种情况 ,而不是如何更改我的代码或如何修复它,因为我也可以做到。 我正在使用 python2.7.3
我正在更新我的包。但是,我正在为依赖项/导入而苦苦挣扎。我使用了两个冲突的包 - ggplot2和 psych及其功能 alpha当然还有 alpha ggplot2 的对象不同于 alpha psy
我是一名优秀的程序员,十分优秀!