代码之家  ›  专栏  ›  技术社区  ›  fractals Francisco Flores

TensorFlow:不同输出形状的数据集之间的交替

  •  2
  • fractals Francisco Flores  · 技术社区  · 8 年前

    tf.Dataset 对于三维图像CNN,其中从训练集和验证集输入的三维图像的形状不同(训练:(64,64,64),验证:(176,176,160)。我甚至不知道这是可能的,但我正在根据一篇论文重新创建这个网络,并使用经典的 feed_dict 方法网络确实有效。出于性能方面的原因(也是为了学习),我正试图将网络切换到 数据集

    我有两个数据集和迭代器,如下所示:

    def _data_parser(dataset, shape):
            features = {"input": tf.FixedLenFeature((), tf.string),
                        "label": tf.FixedLenFeature((), tf.string)}
            parsed_features = tf.parse_single_example(dataset, features)
    
            image = tf.decode_raw(parsed_features["input"], tf.float32)
            image = tf.reshape(image, shape + (1,))
    
            label = tf.decode_raw(parsed_features["label"], tf.float32)
            label = tf.reshape(label, shape + (1,))
            return image, label
    
    train_datasets = ["train.tfrecord"]
    train_dataset = tf.data.TFRecordDataset(train_datasets)
    train_dataset = train_dataset.map(lambda x: _data_parser(x, (64, 64, 64)))
    train_dataset = train_dataset.batch(batch_size) # batch_size = 16
    train_iterator = train_dataset.make_initializable_iterator()
    
    val_datasets = ["validation.tfrecord"]
    val_dataset = tf.data.TFRecordDataset(val_datasets)
    val_dataset = val_dataset.map(lambda x: _data_parser(x, (176, 176, 160)))
    val_dataset = val_dataset.batch(1)
    val_iterator = val_dataset.make_initializable_iterator()
    

    TensorFlow documentation 有关于使用 reinitializable_iterator feedable_iterator ,但它们都在 相同的 输出形状,这里不是这样的。

    我应该如何使用 tf.data.Iterator 对我来说呢?

    1 回复  |  直到 8 年前
        1
  •  3
  •   P-Gn    8 年前

    仅提供未指定的( None )尺寸不匹配的轴上形状的值。例如。

    import numpy as np
    import tensorflow as tf
    
    training_dataset = tf.data.Dataset.from_tensors(np.zeros((64, 64, 64), np.float32)).repeat().batch(4)
    validation_dataset = tf.data.Dataset.from_tensors(np.zeros((176, 176, 160), np.float32)).repeat().batch(1)
    
    iterator = tf.data.Iterator.from_structure(
        training_dataset.output_types,
        tf.TensorShape([None, None, None, None]))
    next_element = iterator.get_next()
    
    training_init_op = iterator.make_initializer(training_dataset)
    validation_init_op = iterator.make_initializer(validation_dataset)
    
    sess = tf.InteractiveSession()
    sess.run(training_init_op)
    print(sess.run(next_element).shape)
    sess.run(validation_init_op)
    print(sess.run(next_element).shape)
    
    
    推荐文章