【发布时间】:2018-07-26 11:34:14
【问题描述】:
我在尝试创建 tf.Dataset 时遇到问题
通过tf.data.TFRecordDataset 来自tfrecord 文件。
def parse_function(example_proto):
# Defaults are not specified since both keys are required.
keys_to_features={
'image': tf.FixedLenFeature([1024*1024],tf.int64),
'label': tf.FixedLenFeature([1024*1024],tf.int64)
}
features = tf.parse_example([example_proto],keys_to_features)
label = features['label']
image = features['image']
label = tf.reshape(label,(1024,1024))
image = tf.reshape(image,(1024,1024))
return image,label
def make_batch(batch_size):
filenames = ["train.tfrecords"]
tf.data.TFRecordDataset(filenames).repeat()
dataset.map(map_func=parse_function,num_parallel_calls=batch_size)
dataset.batch(batch_size)
iterator = dataset.make_one_shot_iterator()
image , label = iterator.get_next()
return image , label
这导致了错误:
当未启用急切执行时,张量对象不可迭代。要迭代此张量,请使用 tf.map_fn。
所以我改变了:image , label = iterator.get_next()
致:next_elem = iterator.get_next()
有了这个我可以执行这个代码:
with tf.Session() as sess:
sess.run(tf.global_variables_initializer())
next_elem = sess.run( make_batch(1))
但是,next_elem 是字节数组,而不是形状为 ([1024,1024],[1024,1024]) 的元组。
【问题讨论】:
标签: python tensorflow tensorflow-datasets tfrecord