You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.27 16:09:23