TensorFlow 2中多GPU数据集分配:解决OOM问题
解决多GPU训练OOM:按GPU分配数据集子集的实现方案
核心思路
使用tf.distribute.MirroredStrategy.distribute_datasets_from_function(),让每个GPU只加载对应的数据子集,而非复制全量数据集。这样双16GB GPU的总可用内存(32GB)就能覆盖单24GB GPU的需求,避免内存溢出。
完整代码示例
1. 准备数据(匹配你的numpy字典输入)
假设输入字典X_dict包含多模态特征,目标字典y_dict包含多回归和概率输出:
import tensorflow as tf import numpy as np # 模拟多模态输入字典 num_samples = 10000 X_dict = { "text_emb": np.random.rand(num_samples, 768).astype(np.float32), # 文本嵌入特征 "transaction_feat": np.random.rand(num_samples, 12).astype(np.float32), # 交易类特征 "structured_feat": np.random.rand(num_samples, 8).astype(np.float32) # 结构化特征 } # 模拟多输出目标字典 y_dict = { "regression_out": np.random.rand(num_samples, 3).astype(np.float32), # 多回归输出 "prob_out": np.random.rand(num_samples, 2).astype(np.float32) # 概率输出 }
2. 定义数据分发函数
该函数为每个GPU分配专属数据集子集:
def dataset_fn(input_context): # 获取当前GPU的ID和总GPU数量 replica_id = input_context.input_pipeline_id num_replicas = input_context.num_replicas_in_sync # 计算当前GPU负责的样本区间 total_samples = num_samples samples_per_replica = total_samples // num_replicas start_idx = replica_id * samples_per_replica end_idx = start_idx + samples_per_replica # 切分输入和目标的子集 X_subset = {key: val[start_idx:end_idx] for key, val in X_dict.items()} y_subset = {key: val[start_idx:end_idx] for key, val in y_dict.items()} # 转换为tf.data.Dataset并添加预处理逻辑 dataset = tf.data.Dataset.from_tensor_slices((X_subset, y_subset)) dataset = dataset.shuffle(samples_per_replica).batch(32) # 替换为你的单GPU batch size return dataset
3. 构建分布式策略与模型
# 初始化MirroredStrategy distributed_strategy = tf.distribute.MirroredStrategy() num_replicas = distributed_strategy.num_replicas_in_sync print(f"使用 {num_replicas} 个GPU训练") # 在策略作用域内构建、编译模型 with distributed_strategy.scope(): def build_model(): # 定义多模态输入层 text_input = tf.keras.layers.Input(shape=(768,), name="text_emb") transaction_input = tf.keras.layers.Input(shape=(12,), name="transaction_feat") structured_input = tf.keras.layers.Input(shape=(8,), name="structured_feat") # 模拟多模态融合逻辑(替换为你的实际模型结构) text_branch = tf.keras.layers.Dense(128, activation="relu")(text_input) transaction_branch = tf.keras.layers.Dense(64, activation="relu")(transaction_input) structured_branch = tf.keras.layers.Dense(64, activation="relu")(structured_input) fused = tf.keras.layers.concatenate([text_branch, transaction_branch, structured_branch]) fused = tf.keras.layers.Dense(256, activation="relu")(fused) # 定义多输出层 regression_out = tf.keras.layers.Dense(3, name="regression_out")(fused) prob_out = tf.keras.layers.Dense(2, activation="sigmoid", name="prob_out")(fused) return tf.keras.Model( inputs=[text_input, transaction_input, structured_input], outputs=[regression_out, prob_out] ) train_model = build_model() # 编译模型(替换为你的损失函数、优化器和指标) train_model.compile( optimizer=tf.keras.optimizers.Adam(), loss={ "regression_out": tf.keras.losses.MeanSquaredError(), "prob_out": tf.keras.losses.BinaryCrossentropy() }, metrics={ "regression_out": ["mae"], "prob_out": ["accuracy"] } )
4. 创建分布式数据集并启动训练
# 生成分布式训练数据集 train_dataset = distributed_strategy.distribute_datasets_from_function(dataset_fn) # 开始训练(无需手动指定steps_per_epoch,TensorFlow会自动计算) train_model.fit( train_dataset, epochs=10 )
关键细节说明
- 内存优化逻辑:每个GPU仅加载
num_samples//num_replicas数量的样本,避免全量数据重复占用GPU内存。 - 预处理适配:所有预处理操作(如标准化、特征工程)都放在
dataset_fn中,确保每个GPU的子集执行相同逻辑。 - batch size调整:代码中
batch(32)是单GPU的batch size,总batch size会自动乘以GPU数量(比如2个GPU时总batch为64),可根据内存灵活调整单GPU的batch大小。
内容的提问来源于stack exchange,提问作者Stefan Meili
相关产品推荐
相关产品推荐

