多GPU使用MirroredStrategy时无法应用AutoShardPolicy.FILE问题咨询
问题根因
FILE自动分片的生效前提是:TensorFlow能沿着tf.data流水线向上溯源,找到最上游持有全量文件/样本列表的可分片数据源,才能把全量样本按文件维度切分给不同GPU加载,避免多GPU重复读盘。
当前流水线最上游使用tf.data.Dataset.from_generator构建,该算子底层对应报错信息中提到的FlatMapDataset节点:生成器的执行逻辑对TensorFlow是黑盒,框架无法获取生成器内部的全量样本列表,也无法感知样本和磁盘文件的对应关系,因此无法执行FILE分片,只能自动回退到DATA分片策略——即每个GPU都加载全量数据集,再按批次切分样本,无法达到预期的IO加速效果。
修复方案
核心修改逻辑是把黑盒的生成器数据源替换为TensorFlow可识别的可分片上游数据源,原有数据读取、预处理逻辑完全不需要改动,和原逻辑100%等价。
- 提前离线遍历
self.list_IDs,把原generator中输出的样本路径、query值、标签整理为三个独立列表,这一步的计算开销和原生成器初始化逻辑完全一致,无额外成本。 - 最上游使用
tf.data.Dataset.from_tensor_slices加载三个列表构建数据集,该算子是TensorFlow原生支持分片溯源的标准数据源,FILE策略可以正常识别。 - 原有的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
相关产品推荐
相关产品推荐

