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

多GPU使用MirroredStrategy时无法应用AutoShardPolicy.FILE问题咨询

问题根因

FILE自动分片的生效前提是:TensorFlow能沿着tf.data流水线向上溯源,找到最上游持有全量文件/样本列表的可分片数据源,才能把全量样本按文件维度切分给不同GPU加载,避免多GPU重复读盘。
当前流水线最上游使用tf.data.Dataset.from_generator构建,该算子底层对应报错信息中提到的FlatMapDataset节点:生成器的执行逻辑对TensorFlow是黑盒,框架无法获取生成器内部的全量样本列表,也无法感知样本和磁盘文件的对应关系,因此无法执行FILE分片,只能自动回退到DATA分片策略——即每个GPU都加载全量数据集,再按批次切分样本,无法达到预期的IO加速效果。

修复方案

核心修改逻辑是把黑盒的生成器数据源替换为TensorFlow可识别的可分片上游数据源,原有数据读取、预处理逻辑完全不需要改动,和原逻辑100%等价。

  1. 提前离线遍历self.list_IDs,把原generator中输出的样本路径、query值、标签整理为三个独立列表,这一步的计算开销和原生成器初始化逻辑完全一致,无额外成本。
  2. 最上游使用tf.data.Dataset.from_tensor_slices加载三个列表构建数据集,该算子是TensorFlow原生支持分片溯源的标准数据源,FILE策略可以正常识别。
  3. 原有的shuffle、map、cache、batch、prefetch、分片策略配置全部保留,不需要调整。

修正后的数据集构建代码如下:

def get_dataset(self):
    AUTOTUNE = tf.data.experimental.AUTOTUNE
    # 提前整理所有样本元信息,替代原生成器的输出逻辑
    all_sample_paths = []
    all_queries = []
    all_labels = []
    for sample in self.list_IDs:
        query_ = int(sample[sample.find('query-') + len('query-'): len(sample)])
        query = self.queries_rescaled[query_] if self.queries_rescaled is not None else query_
        all_sample_paths.append(sample.encode('utf-8'))  # 匹配tf.string类型要求
        all_queries.append(query)
        all_labels.append(self.labels[sample])
    
    # 构建可被FILE分片识别的上游数据源
    dataset = tf.data.Dataset.from_tensor_slices((all_sample_paths, all_queries, all_labels))
    
    if self.shuffle is True:
        dataset = dataset.shuffle(self.num_IDs)
    # 原有map逻辑完全复用,不需要修改
    dataset = dataset.map(self.map_generator, num_parallel_calls=self.num_threads)
    if self.cache is True:
        dataset = dataset.cache(self.cache_path)
    dataset = dataset.batch(self.batch_size, drop_remainder=self.drop_remainder)
    if self.prefetch is True:
        dataset = dataset.prefetch(buffer_size=AUTOTUNE)

    # 原有分片配置保留即可,此时可正常触发FILE分片
    options = tf.data.Options()
    options.experimental_distribute.auto_shard_policy = tf.data.experimental.AutoShardPolicy.AUTO
    dataset = dataset.with_options(options)

    return dataset

def map_generator(self, x_elem, query_elem, label_elem):
    x_input = tf.numpy_function(func=self.get_input, inp=[x_elem], Tout=self.Tout_dtype)
    return {"encoder_input": x_input, "query": query_elem}, label_elem
补充说明
  • 如果后续样本规模扩大,不想把全量元信息存在内存中,可以把所有样本路径、query、标签提前写入TFRecord文件或者文本文件,最上游替换为tf.data.TFRecordDataset或者tf.data.TextLineDataset读取,这两类数据源同样原生支持FILE分片,加速效果一致。
  • tf.numpy_function包裹的读文件逻辑、map/cache/batch/prefetch等常规流水线算子都不会阻断分片溯源,不需要调整。
  • 该修改不会改变原有数据的输出格式、样本顺序,不需要调整后续训练网络的代码。

内容的提问来源于stack exchange,提问作者fruitflyy

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 18:06:08