Python多进程并行写入同一TFRecord文件报数据损坏错误
问题根因
你遇到的DataLossError完全是多进程并发写入同一个文件导致的。tf.io.TFRecordWriter没有内置并发写入的锁机制,多个进程同时持有同一文件的写入句柄、交替写入二进制内容时,会直接打乱TFRecord的固定存储结构,最终读取时必然报数据损坏。
之前用Pool失败的原因也很简单:你试图把TFRecordWriter实例作为参数传给工作进程,这个基于C++实现的文件句柄对象不支持pickle序列化,跨进程传递必然报错。
推荐方案:分片写入后合并(性能最优)
这是工业界处理大规模TFRecord写入的标准方案,没有并发冲突,同时能把CSV读取、数据预处理、序列化这些CPU/IO密集操作全部分摊到多进程,速度提升最明显。
核心逻辑:
- 把待处理的分组均匀切分给N个工作进程
- 每个进程只写自己专属的临时TFRecord分片,文件名用进程序号/ID区分,从根源避免写入冲突
- 所有进程写完后,在主进程把分片顺序拼接成最终的TFRecord文件——TFRecord本身是线性二进制格式,直接拼接合法分片得到的就是完整合法的最终文件,不需要重新解析重写,合并速度极快。
对应修改后的可运行代码:
import os import pandas as pd import multiprocessing import tensorflow as tf TFR_PATH = "./tfr.tfrecord" BANDS = ["B2", "B3","B4","B5","B6","B7","B8","B8A","B11","B12"] # 进程数不要盲目开太大,云存储读取是IO密集操作,4-8个通常是性价比最高的配置 NUM_WORKERS = min(8, multiprocessing.cpu_count()) # 你原有的prepare_df、serialize_example逻辑保持不变即可 def write_tfrecord_shard(shard_path, df_subset, bands): # 每个进程独立创建自己的writer,仅操作专属分片,无任何并发冲突 with tf.io.TFRecordWriter(shard_path) as writer: for _, grp in df_subset: band_data = {b: [] for b in bands} valid_label = None for _, row in grp.iterrows(): try: df = pd.read_csv(row['uri']) except FileNotFoundError: continue df = prepare_df(df, bands) valid_label = row['FS_crop'].encode() for b in bands: band_data[b].append(list(df[b].astype('Int64'))) # 跳过全文件缺失的空分组 if valid_label is None: continue # 长度填充、扁平化逻辑和你原来的实现完全一致 mlen = max([len(j) for j in band_data[bands[0]]]) npx = len(band_data[bands[0]]) flat_band_data = {k: [] for k in band_data} for k,v in band_data.items(): for b in v: flat_band_data[k].extend(b + [0] * int(mlen - len(b))) example_proto = serialize_example(npx, flat_band_data, valid_label) writer.write(example_proto) def merge_shards(shard_paths, final_path): # 顺序拼接所有分片,不需要解析内容 with open(final_path, 'wb') as out_f: for sp in shard_paths: with open(sp, 'rb') as in_f: out_f.write(in_f.read()) os.remove(sp) # 合并完成后删除临时分片 if __name__ == "__main__": # 多进程代码必须放在主入口判断下,否则Windows/macOS会反复递归启动子进程 gqdf = list(qdf.groupby("field_centroid_str")) # 切分任务 chunk_size = len(gqdf) // NUM_WORKERS + 1 shard_paths = [] tasks = [] for shard_idx, start in enumerate(range(0, len(gqdf), chunk_size)): shard_path = f"./tmp_tfr_shard_{shard_idx}.tfrecord" shard_paths.append(shard_path) tasks.append((shard_path, gqdf[start:start+chunk_size], BANDS)) # 多进程并行写分片 with multiprocessing.Pool(NUM_WORKERS) as pool: pool.starmap(write_tfrecord_shard, tasks) # 合并得到最终文件 merge_shards(shard_paths, TFR_PATH) # 验证读取 raw_dataset = tf.data.TFRecordDataset(TFR_PATH) total_records = 0 for _ in raw_dataset: total_records +=1 print(f"写入完成,总有效记录数:{total_records}")
可选方案:多进程序列化+单进程写入
如果你不想生成临时分片,可以把计算和写入完全解耦:
- 工作进程只负责读取CSV、预处理、生成序列化后的example二进制串,通过
multiprocessing.Queue把结果传给主进程 - 主进程单独持有
TFRecordWriter,循环从队列取二进制串写入文件,所有写操作只在单进程执行,完全没有并发冲突
这个方案的缺点是进程间传输二进制数据会有一定IPC开销,性能比分片方案差10%-30%左右,好处是不需要做后续合并。
避坑提醒
- 不要尝试通过加文件锁的方式实现多进程写同一个文件:锁争抢会带来极大的性能开销,几乎抵消多进程带来的速度提升,实际表现远不如分片方案。
- 不要跨进程传递
TFRecordWriter对象,这类基于C++实现的IO对象基本都不支持pickle序列化,强行传递只会报错。 - 如果最终TFRecord需要直接写入云存储,可以先在本地写分片、合并完成后再一次性上传,比直接流式写云存储稳定性更高。
内容的提问来源于stack exchange,提问作者JmeCS
相关产品推荐
相关产品推荐

