【发布时间】:2017-01-24 06:14:15
【问题描述】:
在 Tensorflow 中,我想将多维数组保存到 TFRecord。例如:
[[1, 2, 3], [1, 2], [3, 2, 1]]
由于我要解决的任务是连续的,因此我尝试使用 Tensorflow 的tf.train.SequenceExample(),并且在写入数据时,我成功地将数据写入 TFRecord 文件。但是,当我尝试使用 tf.parse_single_sequence_example 从 TFRecord 文件加载数据时,我遇到了大量神秘错误:
W tensorflow/core/framework/op_kernel.cc:936] Invalid argument: Name: , Key: input_characters, Index: 1. Number of int64 values != expected. values size: 6 but output shape: []
E tensorflow/core/client/tensor_c_api.cc:485] Name: , Key: input_characters, Index: 1. Number of int64 values != expected. values size: 6 but output shape: []
我用来加载数据的函数如下:
def read_and_decode_single_example(filename):
filename_queue = tf.train.string_input_producer([filename],
num_epochs=None)
reader = tf.TFRecordReader()
_, serialized_example = reader.read(filename_queue)
context_features = {
"length": tf.FixedLenFeature([], dtype=tf.int64)
}
sequence_features = {
"input_characters": tf.FixedLenSequenceFeature([], dtype=tf.int64),
"output_characters": tf.FixedLenSequenceFeature([], dtype=tf.int64)
}
context_parsed, sequence_parsed = tf.parse_single_sequence_example(
serialized=serialized_example,
context_features=context_features,
sequence_features=sequence_features
)
context = tf.contrib.learn.run_n(context_parsed, n=1, feed_dict=None)
print context
我用来保存数据的函数在这里:
# http://www.wildml.com/2016/08/rnns-in-tensorflow-a-practical-guide-and-undocumented-features/
def make_example(input_sequence, output_sequence):
"""
Makes a single example from Python lists that follows the
format of tf.train.SequenceExample.
"""
example_sequence = tf.train.SequenceExample()
# 3D length
sequence_length = sum([len(word) for word in input_sequence])
example_sequence.context.feature["length"].int64_list.value.append(sequence_length)
input_characters = example_sequence.feature_lists.feature_list["input_characters"]
output_characters = example_sequence.feature_lists.feature_list["output_characters"]
for input_character, output_character in izip_longest(input_sequence,
output_sequence):
# Extend seems to work, therefore it replaces append.
if input_sequence is not None:
input_characters.feature.add().int64_list.value.extend(input_character)
if output_characters is not None:
output_characters.feature.add().int64_list.value.extend(output_character)
return example_sequence
欢迎任何帮助。
【问题讨论】:
-
嗨,您能提供更多上下文吗?最好提供一个可以实际运行和测试的最小示例,包括如何将数据保存到文件的步骤。
-
您的示例很难理解,如果您编辑示例以包含相关上下文,您将获得更多帮助。例如 - 查看您在代码中添加注释的链接,很明显您生成序列示例的 sn-p 不包含实际写入数据的代码。
标签: python multidimensional-array tensorflow protocol-buffers