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

TensorFlow多模型绑定不同GPU训练仅GPU0占用问题如何解决?

问题核心原因

你仅在模型实例化阶段指定了设备作用域,前向传播、梯度计算、优化器更新等训练阶段的运算没有被纳入对应设备的作用域中,TensorFlow会默认将这部分运算调度到默认设备GPU:0执行,中间生成的张量、梯度缓存自然都会占用GPU:0的显存,就出现了你观察到的只有GPU0显存持续上涨的现象。

解决方案

  • 每个模型的完整训练逻辑(输入张量拷贝、前向传播、损失计算、梯度计算、权重更新)都要包裹到对应GPU的tf.device上下文管理器中
  • 开启TensorFlow显存按需分配,避免框架预占用全卡显存、跨设备调度资源
  • 修正原有代码中梯度调用、优化器更新的语法错误

修正后代码示例

# 第一步:先配置显存按需分配,防止框架预占全卡显存
gpus = tf.config.list_physical_devices('GPU')
for gpu in gpus:
    tf.config.experimental.set_memory_growth(gpu, True)

# 模型、优化器都在对应设备作用域内初始化,保证变量存储在指定GPU
with tf.device('/device:GPU:0'):
    model1 = model1Class().model()
optimizer1 = tf.keras.optimizers.Adam()

with tf.device('/device:GPU:1'):
    model2 = model2Class().model()
optimizer2 = tf.keras.optimizers.Adam()

with tf.device('/device:GPU:2'):
    model3 = model3Class().model()
optimizer3 = tf.keras.optimizers.Adam()

lossFunc = tf.keras.losses.YourLoss()

for epoch in range(10):
    dataGen = DataGenerator(...)
    X_raw, y = next(dataGen)

    # 模型1完整训练逻辑都套GPU0作用域
    with tf.device('/device:GPU:0'):
        X = tf.identity(X_raw) # 将输入张量拷贝到GPU0
        with tf.GradientTape() as tape1:
            X = model1(X)
            loss1 = lossFunc(X, y[1])
        grads1 = tape1.gradient(loss1, model1.trainable_weights)
        optimizer1.apply_gradients(zip(grads1, model1.trainable_weights))
        # 输出转为CPU张量,避免跨设备直接传输带来的异常调度
        X_out1 = tf.identity(X).cpu()

    # 模型2完整训练逻辑都套GPU1作用域
    with tf.device('/device:GPU:1'):
        X = tf.identity(X_out1) # 将上一个模型的输出拷贝到GPU1
        with tf.GradientTape() as tape2:
            X = model2(X)
            loss2 = lossFunc(X, y[2])
        grads2 = tape2.gradient(loss2, model2.trainable_weights)
        optimizer2.apply_gradients(zip(grads2, model2.trainable_weights))
        X_out2 = tf.identity(X).cpu()

    # 模型3完整训练逻辑都套GPU2作用域
    with tf.device('/device:GPU:2'):
        X = tf.identity(X_out2) # 将上一个模型的输出拷贝到GPU2
        with tf.GradientTape() as tape3:
            X = model3(X)
            loss3 = lossFunc(X, y[3])
        grads3 = tape3.gradient(loss3, model3.trainable_weights)
        optimizer3.apply_gradients(zip(grads3, model3.trainable_weights))

可选排查手段

如果修改后依然有设备调度异常,可以添加tf.debugging.set_log_device_placement(True)代码,运行时会打印每个运算实际运行的设备,方便定位调度错误的节点。

内容的提问来源于stack exchange,提问作者D. Ramsook

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.26 06:15:07