Skip to content

TensorFlow - 分布式计算

TensorFlow - 使用 tf.distribute.Strategy 进行分布式计算入门

Section titled “TensorFlow - 使用 tf.distribute.Strategy 进行分布式计算入门”

分布式训练允许您在多个设备(CPU、GPU、TPU)甚至多台机器上训练机器学习模型。这可以显著加快大型模型或大型数据集的训练速度。TensorFlow 提供了 tf.distribute.Strategy API 作为分布式训练的主要机制。

本章介绍分布式计算的概念,并提供了如何设置多 worker 分布式训练场景的高级概述。较旧的 tf.train.ClusterSpec 和 tf.train.Server API 已弃用;tf.distribute.Strategy 是现代且推荐的方法。

tf.distribute.Strategy 是一个处理在各种硬件配置上分发模型训练复杂性的 API。关键策略包括:

  • MirroredStrategy:支持在单台机器上的多 GPU 上进行同步分布式训练。模型变量在所有 GPU 上镜像,并同步应用更新。
  • MultiWorkerMirroredStrategy:支持跨多个 worker(机器)进行同步分布式训练,每个 worker 可能包含多个 GPU。类似于 MirroredStrategy,但适用于多机器设置。
  • TPUStrategy:用于在 Tensor Processing Units (TPUs) 上进行训练。
  • ParameterServerStrategy:支持异步参数服务器训练,其中一些 worker 充当参数服务器(存储变量),其他 worker 充当计算 worker。

使用这些策略通常只需要对现有 TensorFlow Keras 训练代码进行最少的修改。主要步骤是:

  1. 实例化一个策略。
  2. 将模型创建和编译放在策略的 scope 内(with strategy.scope():)。
  3. 使用 tf.data.Dataset 准备数据集,并确保必要时正确进行分片或分发。

概念示例:MultiWorkerMirroredStrategy

Section titled “概念示例:MultiWorkerMirroredStrategy”

设置 MultiWorkerMirroredStrategy 需要在参与集群的每台机器(worker)上配置 TF_CONFIG 环境变量。此变量告知 TensorFlow 集群的结构(worker 地址、当前 worker 的类型和索引)。

步骤 1:定义 TF_CONFIG(概念 - 在每个 worker 机器上)

假设您有两台 worker 机器,worker0.example.com 和 worker1.example.com,它们都在端口 12345 上监听。

在 worker0.example.com 上,您将设置 TF_CONFIG(例如,在您的 shell 或 Python 脚本中):

import json
import os
# Example TF_CONFIG for worker 0
tf_config_worker0 = {
'cluster': {
'worker': ['worker0.example.com:12345', 'worker1.example.com:12345']
},
'task': {'type': 'worker', 'index': 0}
}
# In a real scenario, this would be set as an environment variable BEFORE the script runs.
# For demonstration, we might set it in Python if this script IS worker 0.
# os.environ['TF_CONFIG'] = json.dumps(tf_config_worker0)

在 worker1.example.com 上,TF_CONFIG 将类似,但 'index': 1。

# Example TF_CONFIG for worker 1
tf_config_worker1 = {
'cluster': {
'worker': ['worker0.example.com:12345', 'worker1.example.com:12345']
},
'task': {'type': 'worker', 'index': 1}
}
# os.environ['TF_CONFIG'] = json.dumps(tf_config_worker1)

步骤 2:用于训练的 Python 脚本(在所有 worker 上运行)

设置 TF_CONFIG 后,将在所有 worker 机器上运行相同的 Python 脚本。

import tensorflow as tf
from tensorflow import keras
from tensorflow.keras import layers
import numpy as np
import os
import json
# --- This part is for simulation if not running in a real multi-worker env ---
# In a real multi-worker setup, TF_CONFIG is an environment variable.
# To simulate locally for this tutorial, we can try to set it if not present.
if 'TF_CONFIG' not in os.environ:
# Simulate a single worker setup if no TF_CONFIG is found for demonstration purposes.
# For true multi-worker, TF_CONFIG must be set externally for each worker.
print("TF_CONFIG not set. Simulating single-worker strategy or defaulting.")
# strategy = tf.distribute.get_strategy() # Gets default strategy
strategy = tf.distribute.MirroredStrategy() # For local multi-GPU, or fallback
if not tf.config.list_physical_devices('GPU'):
print("No GPUs found, MirroredStrategy will run on CPU. May be slow.")
else:
# If TF_CONFIG is set, MultiWorkerMirroredStrategy will be used.
strategy = tf.distribute.MultiWorkerMirroredStrategy()
# You can access cluster information if needed:
# cluster_resolver = tf.distribute.cluster_resolver.TFConfigClusterResolver()
# print(f"Task type: {cluster_resolver.task_type}, Task ID: {cluster_resolver.task_id}")
# print(f"Cluster spec: {cluster_resolver.cluster_spec()}")
print(f'Number of devices: {strategy.num_replicas_in_sync}')
# Hyperparameters
BUFFER_SIZE = 10000
BATCH_SIZE_PER_REPLICA = 64
GLOBAL_BATCH_SIZE = BATCH_SIZE_PER_REPLICA * strategy.num_replicas_in_sync
EPOCHS = 3
# Create a simple dataset (e.g., MNIST)
(x_train, y_train), _ = keras.datasets.mnist.load_data()
x_train = x_train.astype('float32') / 255.0
x_train = np.expand_dims(x_train, axis=-1)
y_train = keras.utils.to_categorical(y_train, num_classes=10)
train_dataset = tf.data.Dataset.from_tensor_slices((x_train, y_train))
train_dataset = train_dataset.shuffle(BUFFER_SIZE).batch(GLOBAL_BATCH_SIZE)
# Model creation and compilation within strategy.scope()
with strategy.scope():
model = keras.Sequential([
layers.Conv2D(32, kernel_size=(3, 3), activation='relu', input_shape=(28, 28, 1)),
layers.MaxPooling2D(pool_size=(2, 2)),
layers.Flatten(),
layers.Dense(10, activation='softmax')
])
model.compile(optimizer='adam',
loss='categorical_crossentropy',
metrics=['accuracy'])
print("Model created and compiled within strategy scope.")
# List devices visible to TensorFlow within this strategy
# In a real multi-worker setup, this would show devices across the cluster from the perspective
# of the current worker and strategy. For a local MirroredStrategy, it lists local GPUs/CPUs.
# For TF1.x style device listing (sess.list_devices()), the closest in TF2 is checking
# physical devices and how the strategy maps to them.
print("\nPhysical Devices available to TensorFlow:")
for device_type in ['CPU', 'GPU']:
devices = tf.config.list_physical_devices(device_type)
if devices:
print(f"{device_type}s: {devices}")
else:
print(f"No {device_type}s found.")
# Train the model
if strategy.num_replicas_in_sync > 0: # Ensure we have replicas
print(f"\nStarting training on {strategy.num_replicas_in_sync} replicas...")
model.fit(train_dataset, epochs=EPOCHS)
print("Training complete.")
else:
print("No replicas available for training. Check strategy setup.")
# For more details on distributed training:
# https://www.tensorflow.org/guide/distributed_training

MultiWorkerMirroredStrategy 的关键点:

  • TF_CONFIG: 对于 worker 之间的相互发现和协调至关重要。
  • 策略 Scope: 模型定义、优化器创建和 model.compile() 必须发生在 with strategy.scope(): 内部。
  • 数据分片 (Sharding): 当使用 distribute_datasets_from_function 或数据集已正确批处理时,tf.data.Dataset 会自动为大多数策略处理数据分片。
  • 全局 Batch Size: 在 dataset.batch() 中指定的 batch size 应该是全局 batch size(每个 replica 的 batch size * replica 数量)。
  • 保存/加载模型: 使用 tf.keras.callbacks.ModelCheckpoint 和 tf.keras.models.load_model。对于 MultiWorkerMirroredStrategy,建议只有主 worker(worker 0)保存 checkpoint 以避免冲突。这可以通过 tf.distribute.experimental.coordinator.ClusterCoordinator 或自定义 callback 来管理。

此示例提供了概念性概述。设置真正的多 worker 环境需要网络配置并确保所有 worker 都能通信。与旧的、更手动的方法相比,tf.distribute.Strategy API 极大地简化了编写分布式训练代码的过程。