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

如何在多GPU环境下将静态张量分发至各GPU并在tf.data中使用本地副本

多GPU环境下tf.data静态数组设备内存绑定问题

我有一个tf.data.Dataset,通过整数索引列表从静态数组中选取数据。为了避免主机-设备数据拷贝占用50%训练时间,静态数组需要存储在GPU内存中(该数组可轻松容纳于GPU内存)。单GPU下的工作代码如下:

host_array = np.concatenate([self.cache[filename] for filename in self.filenames])
with tf.device('/gpu:0'):
    device_array = tf.constant(host_array)
    frag_length = tf.constant([self.fragment_length])
    ds = tf.data.Dataset.from_tensor_slices(index_list)
    @tf.function
    def get_segment_tuple(index):
        sample = tf.slice(device_array, [index], frag_length)
        return (sample, sample)

    if self.shuffle:
        ds = ds.shuffle(len(index_list))

    x_ds = ds.map(get_segment_tuple, num_parallel_calls=tf.data.AUTOTUNE)
    return x_ds.batch(self.batch_size)

希望适配多GPU机器,使用tf.keras.utils.experimental.DatasetCreator传入model.fit,但不知道在dataset_fn里的tf.device应该填什么——InputContext没有device属性,需要明确将静态数组放到对应GPU,确保切片操作使用本地设备内存,避免跨设备/主机-设备拷贝。同时还有疑问:每个worker的shuffle是否会产生重复结果?


解决方案

核心思路

在单机多GPU训练中,input_context.input_pipeline_id对应每个GPU副本的索引(从0开始),可以用这个ID构造目标GPU设备路径,将静态数组绑定到对应GPU内存;同时通过为每个worker分配唯一随机种子,解决shuffle结果重复的问题。

修改后的完整代码

1. 重构数据集生成函数

def make_local_ds(host_array, frag_length, index_list, shuffle=False, shuffle_seed=None):
    # 在指定GPU内存中创建静态数组
    device_array = tf.constant(host_array)
    frag_length_tensor = tf.constant([frag_length])
    ds = tf.data.Dataset.from_tensor_slices(index_list)
    
    @tf.function
    def get_segment_tuple(index):
        # 直接访问本地GPU内存中的数组,无跨设备拷贝
        sample = tf.slice(device_array, [index], frag_length_tensor)
        return (sample, sample)

    if shuffle:
        # 每个worker使用不同种子,避免shuffle结果重复
        ds = ds.shuffle(len(index_list), seed=shuffle_seed)

    return ds.map(get_segment_tuple, num_parallel_calls=tf.data.AUTOTUNE)

2. 分布式数据集构造函数

def dataset_fn(input_context):
    # 通过input_pipeline_id获取当前worker对应的GPU设备
    device_id = input_context.input_pipeline_id
    target_device = f'/gpu:{device_id}'
    
    with tf.device(target_device):
        # 生成唯一shuffle种子
        shuffle_seed = 42 + input_context.input_pipeline_id
        dataset = make_local_ds(
            host_array, 
            self.fragment_length, 
            index_list, 
            shuffle=self.shuffle, 
            shuffle_seed=shuffle_seed
        )
    
    batch_size = input_context.get_per_replica_batch_size(global_batch_size)
    dataset = dataset.shard(input_context.num_input_pipelines, input_context.input_pipeline_id)
    return dataset.batch(batch_size)

关键说明

  • 设备绑定逻辑:input_pipeline_id与GPU设备索引一一对应(单机多GPU环境下),通过f'/gpu:{device_id}'显式指定设备,确保静态数组存储在当前worker对应的GPU内存中,切片操作直接访问本地设备内存,彻底避免跨设备或主机-设备拷贝。
  • Shuffle去重:基于input_pipeline_id生成唯一随机种子,保证每个worker的数据集shuffle结果不同,避免训练数据重复。
  • 分布式策略适配:需配合tf.distribute.MirroredStrategy等单机多GPU策略使用,TensorFlow会自动将每个输入流水线调度到对应GPU,与显式设备绑定逻辑互补,确保性能最优。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.23 07:15:51