【问题标题】:Unable to batch dataset using `.batch` and `.padded_batch`无法使用“.batch”和“.padded_batch”批处理数据集
【发布时间】:2022-08-22 14:34:33
【问题描述】:

我正在向 tfrecord 写入一些可变长度字符串功能。如果所有示例的特征都具有相同的形状,则它运行得非常好,没有问题。如果形状不同,则每当读取创建的 tfrecord 时都会引发以下错误。

import random

import numpy as np
import tensorflow as tf


def serialize_example(writer):
    # s = np.array([\'aaa\' for _ in range(10)])  # this works fine
    s = np.array([\'aaa\' for _ in range(random.randint(1, 100))])
    features = {
        \'f1\': tf.train.Feature(
            bytes_list=tf.train.BytesList(value=[tf.io.serialize_tensor(s).numpy()])
        )
    }
    example = tf.train.Example(features=tf.train.Features(feature=features))
    writer.write(example.SerializeToString())


def create_tfrecord(output_path):
    with tf.io.TFRecordWriter(output_path) as writer:
        for i in range(total := 100):
            print(f\'\\rWriting example: {i + 1}/{total}\', end=\'\')
            serialize_example(writer)


def read_example(example, feature_map):
    features = tf.io.parse_single_example(example, feature_map)
    f1 = tf.sparse.to_dense(features[\'f1\'])
    f1 = tf.io.parse_tensor(f1[0], tf.string)
    return f1


def read_tfrecord(fp, batch_size):
    files = tf.data.Dataset.list_files(fp)
    dataset = files.flat_map(tf.data.TFRecordDataset)
    feature_map = {
        \'f1\': tf.io.VarLenFeature(tf.string),
    }
    return dataset.map(
        lambda x: read_example(x, feature_map),
        tf.data.experimental.AUTOTUNE,
    ).batch(batch_size)  # if this is removed, both cases work fine


if __name__ == \'__main__\':
    create_tfrecord(\'xyz.tfrecord\')
    dataset = read_tfrecord(\'xyz.tfrecord\', 8)
    sample = dataset.take(1).as_numpy_iterator().next()

错误:

tensorflow.python.framework.errors_impl.InvalidArgumentError: Cannot add tensor to the batch: number of elements does not match. Shapes are: [tensor]: [83], [batch]: [32] [Op:IteratorGetNext]

如果 .batch(batch_size) 被删除,它在这两种情况下都可以正常工作。我期望用.padded_batch(batch_size) 替换.batch 可以解决问题,但是,由于tensorflow 的出色实现会产生未知的形状,这也是完全不可能的。

ValueError: You must provide `padded_shapes` argument because component 0 has unknown rank.

当然,不可能知道read_example 中缺少的padded_shapes

    标签: python tensorflow tf.data.dataset


    【解决方案1】:

    由此,我们可以看出张量的形状不同

    if __name__ == '__main__':
        create_tfrecord('xyz.tfrecord')
        dataset = read_tfrecord('xyz.tfrecord', 8)
        for raw_record in dataset.take(2):
          print(raw_record)
        sample = dataset.take(1).as_numpy_iterator().next()
    

    输出

    Writing example: 100/100tf.Tensor(
    [b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa'
     b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa'
     b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa'
     b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa'
     b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa'
     b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa'
     b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa'
     b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa'
     b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa'
     b'aaa' b'aaa' b'aaa'], shape=(93,), dtype=string)
    tf.Tensor(
    [b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa'
     b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa'
     b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa' b'aaa'
     b'aaa' b'aaa' b'aaa' b'aaa' b'aaa'], shape=(35,), dtype=string)
    

    要批量处理这些张量,我们应该先取消批量处理,然后再批量处理。请在下面的代码中添加.apply(tf.data.experimental.unbatch()),如下所示。

    def read_tfrecord(fp, batch_size):
        files = tf.data.Dataset.list_files(fp)
        dataset = files.flat_map(tf.data.TFRecordDataset)
        feature_map = {
            'f1': tf.io.VarLenFeature(tf.string),
        }
        return dataset.map(
            lambda x: read_example(x, feature_map),
            tf.data.experimental.AUTOTUNE,
            
        ).apply(tf.data.experimental.unbatch()).batch(batch_size)  #unbatch and then batch
    

    请找到完整代码here

    【讨论】:

      猜你喜欢
      • 1970-01-01
      • 2020-04-23
      • 1970-01-01
      • 1970-01-01
      • 2020-08-06
      • 1970-01-01
      • 2020-05-29
      • 2017-05-24
      • 2020-10-16
      相关资源
      最近更新 更多