【问题标题】:Dataset with variable-sized items from tf.dynamic_partition具有来自 tf.dynamic_partition 的可变大小项目的数据集
【发布时间】:2018-07-17 00:58:59
【问题描述】:

类似于this question,我想从一个包含不同大小的每个元素的列表中构建一个TF dataset。但是,与链接的问题不同,我想从tf.dynamic_partition 的输出生成数据集,它输出张量列表。

我的设置:

import tensorflow as tf
D = tf.data.Dataset # shorthand notation

x = tf.range(9) # Array to be partitioned
p = tf.constant([1,0,2,0,0,0,2,2,1]) # Defines partitions

因此,数据集应包含三个元素,分别包含 [1 3 4 5]、[0 8] 和 [2 6 7]。

正如预期的那样,直接方法失败了:

dataset = D.from_tensor_slices(tf.dynamic_partition(x,p,3))
iterator = dataset.make_one_shot_iterator()
next_element = iterator.get_next()
with tf.Session() as sess:
    nl = sess.run(next_element)

tensorflow.python.framework.errors_impl.InvalidArgumentError:形状 所有输入必须匹配: values[0].shape = [4] != values[1].shape = [2]

接下来我尝试的是应用solution of the linked question,应用from_generator:

dataset = D.from_generator(lambda: tf.dynamic_partition(x,p,3), tf.int32, output_shapes=[None])
iterator = dataset.make_one_shot_iterator()
next_element = iterator.get_next()
with tf.Session() as sess:
    nl = sess.run(next_element)

tensorflow.python.framework.errors_impl.InvalidArgumentError: exceptions.ValueError: 使用序列设置数组元素。

如何从tf.dynamic_partition 的输出中创建包含可变大小项目的数据集?

【问题讨论】:

    标签: python tensorflow tensorflow-datasets


    【解决方案1】:

    from_generator 不起作用,因为它希望生成器函数生成 numpy 数组而不是张量。

    解决问题的一种方法是为分区的每个元素创建一个数据集。在您的情况下,您将数据划分为 3 个组,因此您将创建 3 个数据集并将它们与 tf.data.Dataset.concatenate() 组合:

    x = tf.range(9)  # Array to be partitioned
    p = tf.constant([1, 0, 2, 0, 0, 0, 2, 2, 1])  # Defines partitions
    
    partition = tf.dynamic_partition(x, p, 3)
    
    dataset = tf.data.Dataset.from_tensors(partition[0])
    for i in range(1, 3):
        dataset_bis = tf.data.Dataset.from_tensors(partition[i])
        dataset = dataset.concatenate(dataset_bis)
    
    iterator = dataset.make_one_shot_iterator()
    next_element = iterator.get_next()
    
    
    with tf.Session() as sess:
        for i in range(3):
            nl = sess.run(next_element)
            print(nl)
    

    【讨论】:

    • 这行得通!稍微有点之外,但我想提一下,如果有多个对concatenate 的调用,这似乎确实会严重放慢速度。我尝试将第一个循环的范围增加到range(1,20) 并与dataset_bis = tf.data.Dataset.from_tensors(partition[i%3]) 连接。提供数据确实非常非常慢。根据您拥有的分区数量,可能可行或不可行。
    • @mikkola:是的,我想这真的取决于你项目的细节。对于您的数据集,可能会有更好的解决方案。
    猜你喜欢
    • 2020-07-07
    • 2019-02-06
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2019-03-02
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多