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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 06:55:21