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

TensorFlow分布式训练为何仅用单服务器?如何强制使用多台?

问题

我按照TensorFlow分布式训练文档进行训练,执行以下代码:

strategy = tf.distribute.MirroredStrategy()
print('Number of devices: {}'.format(strategy.num_replicas_in_sync))

输出结果为2,但在Databricks Ganglia中仅显示1台服务器被使用。当前有2台可用服务器,请问出现该情况的原因是什么?是否有办法强制让多台服务器分摊训练任务?

相关模型构建代码如下:

with strategy.scope():
  model = tf.keras.Sequential([
    tf.keras.layers.Conv2D(32, 3, activation='relu', input_shape=(28, 28, 1)),
    tf.keras.layers.MaxPooling2D(),
    tf.keras.layers.Flatten(),
    tf.keras.layers.Dense(64, activation='relu'),
    tf.keras.layers.Dense(10)
  ])

  model.compile(loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),
            optimizer=tf.keras.optimizers.Adam(),
            metrics=['accuracy'])

原因分析

  • MirroredStrategy的本质限制:MirroredStrategy是单节点多设备的分布式策略,仅能利用同一台服务器上的多个GPU/CPU,无法跨多台服务器分配任务。你看到的num_replicas_in_sync=2,实际是当前单台服务器上有2个可用计算设备(比如2块GPU),并非2台独立服务器。
  • 集群配置未适配跨节点训练:默认情况下,若未指定跨节点分布式策略,TensorFlow只会使用当前驱动节点的设备资源,不会自动调度到其他服务器。

解决方案

要实现多台服务器分摊训练任务,需改用MultiWorkerMirroredStrategy——这是TensorFlow专门针对多节点分布式训练设计的策略,具体步骤如下:

1. 配置集群拓扑信息

在Databricks中,需通过TF_CONFIG环境变量告知TensorFlow集群的节点分布,示例配置(可在Notebook中执行):

import json
import os

os.environ['TF_CONFIG'] = json.dumps({
    'cluster': {
        'worker': ['worker-node-1:2222', 'worker-node-2:2222']
    },
    'task': {'type': 'worker', 'index': 0}
})

注意:worker列表需替换为你集群中各节点的实际地址,task.index对应当前节点的索引(驱动节点一般设为0,worker节点依次递增)。

2. 替换分布式策略

将原代码中的MirroredStrategy替换为MultiWorkerMirroredStrategy:

strategy = tf.distribute.MultiWorkerMirroredStrategy()
print('Number of devices: {}'.format(strategy.num_replicas_in_sync))

3. 保持模型构建逻辑不变

你现有的模型定义、编译代码已经正确放在strategy.scope()作用域内,无需修改,确保所有与模型相关的操作都在该作用域中执行。

4. 简化配置的替代方案(Databricks专属)

可以使用Databricks提供的Spark TensorFlow Distributor自动处理集群调度和TF_CONFIG配置,示例代码:

from spark_tensorflow_distributor import MirroredStrategyRunner

def train_fn():
    strategy = tf.distribute.MultiWorkerMirroredStrategy()
    with strategy.scope():
        # 复用你原有的模型构建与编译代码
        model = tf.keras.Sequential([
            tf.keras.layers.Conv2D(32, 3, activation='relu', input_shape=(28, 28, 1)),
            tf.keras.layers.MaxPooling2D(),
            tf.keras.layers.Flatten(),
            tf.keras.layers.Dense(64, activation='relu'),
            tf.keras.layers.Dense(10)
        ])
        model.compile(loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),
                      optimizer=tf.keras.optimizers.Adam(),
                      metrics=['accuracy'])
    # 加载数据并启动训练
    (x_train, y_train), _ = tf.keras.datasets.mnist.load_data()
    x_train = x_train.reshape(-1, 28, 28, 1).astype('float32') / 255.0
    model.fit(x_train, y_train, epochs=5, batch_size=64)

# 指定使用2个节点进行训练
runner = MirroredStrategyRunner(num_slots=2, local_mode=False)
runner.run(train_fn)

完成上述配置后,启动训练即可看到多台服务器的资源被同时占用。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.13 03:10:26