【问题标题】:TensorFlow - Error when using interleave or parallel_interleaveTensorFlow - 使用 interleave 或 parallel_interleave 时出错
【发布时间】:2019-07-15 18:21:30
【问题描述】:

我正在使用像 Q&A 这样的 V1.12 API 的 tf.data.Datasets 来读取目录中每个文件预先保存的多个 .h5 文件。 我先做了一个生成器:

class generator_yield:
    def __init__(self, file):
        self.file = file

    def __call__(self):
        with h5py.File(self.file, 'r') as f:
            yield f['X'][:], f['y'][:]

然后制作一个文件名列表并将它们传递给Dataset

def _fnamesmaker(dir, mode='h5'):
    fnames = []
    for dirpath, _, filenames in os.walk(dir):
        for fname in filenames:
            if fname.endswith(mode):
                fnames.append(os.path.abspath(os.path.join(dirpath, fname)))
    return fnames

fnames = _fnamesmaker('./')
len_fnames = len(fnames)
fnames = tf.data.Dataset.from_tensor_slices(fnames)

应用Dataset的interleave方法:

# handle multiple files
ds = fnames.interleave(lambda filename: tf.data.Dataset.from_generator(
    generator_yield(filename), output_types=(tf.float32, tf.float32),
    output_shapes=(tf.TensorShape([100, 100, 1]), tf.TensorShape([100, 100, 1]))), cycle_length=len_fnames)
ds = ds.batch(5).shuffle(5).prefetch(5)

# init iterator
it = ds.make_initializable_iterator()
init_op = it.initializer
X_it, y_it = it.get_next()

型号:

# model
with tf.name_scope("Conv1"):
    W = tf.get_variable("W", shape=[3, 3, 1, 1],
                         initializer=tf.contrib.layers.xavier_initializer())
    b = tf.get_variable("b", shape=[1], initializer=tf.contrib.layers.xavier_initializer())
    layer1 = tf.nn.conv2d(X_it, W, strides=[1, 1, 1, 1], padding='SAME') + b
    logits = tf.nn.relu(layer1)


    loss = tf.reduce_mean(tf.losses.mean_squared_error(labels=y_it, predictions=logits))
    train_op = tf.train.AdamOptimizer(learning_rate=0.0001).minimize(loss)

开始会话:

with tf.Session() as sess:
    sess.run([tf.global_variables_initializer(), init_op])
    while True:
        try:
            data = sess.run(train_op)
            print(data.shape)
        except tf.errors.OutOfRangeError:
            print('done.')
            break

错误看起来像:

TypeError: 预期的 str、bytes 或 os.PathLike 对象,而不是 Tensor 在生成器的 init 方法中。显然,当一个应用交错时,它是一个张量传递到生成器

【问题讨论】:

    标签: tensorflow tensorflow-datasets h5py


    【解决方案1】:

    您不能直接通过 sess.run 运行数据集对象。您必须定义一个迭代器,获取下一个元素。尝试做类似的事情:

    next_elem = files.make_one_shot_iterator.get_next()
    data = sess.run(next_elem)
    

    你应该能够得到你的张量。

    【讨论】:

    • @Zeliang Su 所有数据集对象都需要使用一些迭代器来获取元素。所有方法都是转换或聚合,并且这些结果只能通过迭代器访问。查看Importing data 上的指南,更清楚地了解 API 的结构。
    【解决方案2】:

    根据这个post,我的情况不会从parralel_interleave 的性能中受益。

    ...具有转换源的每个元素的转换 将数据集分成多个元素到目标数据集中...

    它在典型的分类问题中更相关,数据(狗、猫...)保存在单独的目录中。我们这里有一个分割问题,这意味着标签包含输入图像的相同维度。所有数据都存储在一个目录中,每个 .h5 文件都包含一个图像及其标签(掩码)

    这里,一个简单的mapnum_parallel_callssufficient

    【讨论】:

      猜你喜欢
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2020-09-11
      • 2018-10-22
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      相关资源
      最近更新 更多