如何在多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
相关产品推荐
相关产品推荐

