由于您正在使用 tensorflow 中的对象检测,因此官方 Tensorflow models 存储库中有一些不错的代码可以满足您的需求。请注意,此代码适用于 Tensorflow2(不确定它是否适用于 TF1)
参见 example 从 coco 注释中编写分片 tfrecords。这个想法是您在退出堆栈中打开一个 TFRecordWriter 列表(使用contextlib2.ExitStack()),当每个线程完成写入时,它将自动关闭 TFRecords。
实用函数open_sharded_output_tfrecords 函数创建这个 TFRecordWriter 列表
import contextlib2
import tensorflow as tf
with contextlib2.ExitStack() as tf_record_close_stack, tf.gfile.GFile(
annotations_file, 'r'
) as fid:
output_tfrecords = tf_record_creation_util.open_sharded_output_tfrecords(
tf_record_close_stack, output_path, num_shards
)
接下来,您可以使用 ProcessPoolExecutor 以循环方式并行将 tfrecords 写入每个分片(本例中为 4 个工作人员)
from concurrent.futures.process import ProcesPoolExecutor
with ProcessPoolExecutor(4) as executor:
for idx, image in enumerate(images):
futures = []
future = executor.submit(
_write_tf_record,
image,
idx,
num_shards,
output_tfrecords,
)
futures.append(future)
for future in futures:
future.result()
_write_tf_record 可能看起来像这样:
def _write_tf_record(image, idx, num_shards, output_tfrecords)
tf_example = create_tf_example(image)
shard_idx = idx % num_shards
output_tfrecords[shard_idx].write(tf_example.SerializeToString())
只要确保你的分片比多进程工作者多,否则同一个 writer 可能会被两个不同的进程访问。