如何批量写入TFRecords?

Ins*_*ous 7 python file-writing tensorflow

我有一个大约有4000万行的CSV。每行都是一个训练实例。根据有关使用TFRecords的文档,我正在尝试将数据编码并保存在TFRecord文件中。

我发现的所有示例(甚至是TensorFlow repo中的示例)都表明创建TFRecord的过程取决于TFRecordWriter类。此类具有一个方法write,该方法将数据的序列化字符串表示形式输入并将其写入磁盘。但是,这似乎是一次完成一个训练实例。

如何编写一批序列化数据?

假设我有一个功能:

  def write_row(sentiment, text, encoded):
    feature = {"one_hot": _float_feature(encoded),
               "label": _int64_feature([sentiment]),
               "text": _bytes_feature([text.encode()])}

    example = tf.train.Example(features=tf.train.Features(feature=feature))
    writer.write(example.SerializeToString())
Run Code Online (Sandbox Code Playgroud)

写入磁盘4000万次(每个示例一次)将非常慢。批处理此数据并一次写入50k或100k示例(在机器资源允许的范围内)将更加有效。但是,似乎没有任何方法可以在内部执行此操作TFRecordWriter

类似于以下内容:

class MyRecordWriter:

  def __init__(self, writer):
    self.records = []
    self.counter = 0
    self.writer = writer

  def write_row_batched(self, sentiment, text, encoded):
    feature = {"one_hot": _float_feature(encoded),
               "label": _int64_feature([sentiment]),
               "text": _bytes_feature([text.encode()])}

    example = tf.train.Example(features=tf.train.Features(feature=feature))
    self.records.append(example.SerializeToString())
    self.counter += 1
    if self.counter >= 10000:
      self.writer.write(os.linesep.join(self.records))
      self.counter = 0
      self.records = []
Run Code Online (Sandbox Code Playgroud)

但是当读取通过此方法创建的文件时,出现以下错误:

tensorflow/core/framework/op_kernel.cc:1192] Invalid argument: Could not parse example input, value: '
??

label

??
one_hot????
??
Run Code Online (Sandbox Code Playgroud)

注意:我可以更改编码过程,以便每个example原型包含数千个示例,而不仅仅是一个示例,但是我不希望在以这种方式写入TFrecord文件时预先批处理数据,因为这会给我的培训带来额外的开销我想使用该文件进行不同批处理量的训练时使用管道。

de1*_*de1 5

TFRecords是二进制格式。在下面的行中,您将其视为文本文件:self.writer.write(os.linesep.join(self.records))

那是因为您使用的操作系统取决于linesep\n\r\n)。

解决方案:只写记录。您要求批量编写它们。您可以使用缓冲的编写器。对于4000万行,您可能还需要考虑将数据拆分为单独的文件,以实现更好的并行化。

使用TFRecordWriter时:该文件已被缓冲。

在来源中可以找到有关的证据:

  • tf_record.py调用pywrap_tensorflow.PyRecordWriter_New
  • PyRecordWriter调用Env::Default()->NewWritableFile
  • NewWritableFile匹配文件系统上的Env-> NewWritableFile调用
  • 例如PosixFileSystem调用fopen
  • fopen返回的流“如果已知不引用交互式设备,则默认情况下已完全缓冲”
  • 这将取决于文件系统,但WritableFile指出“该实现必须提供缓冲,因为调用者可能一次将小片段附加到文件中。”

  • @ de1您能否举一个简单的示例,说明如何使用`TFRecordWriter`进行批处理? (2认同)