Класс ParameterServerStrategy
Класс ParameterServerStrategy представляет собой стратегию распределенного обучения,
в которой вычисления организуются между несколькими рабочими узлами и параметрическими серверами.
В этой архитектуре параметрические серверы хранят переменные модели, а рабочие узлы выполняют
прямой и обратный проходы, обмениваясь градиентами с серверами. Класс наследуется от
tf.distribute.Strategy и предназначен для использования в многоузловых кластерах.
Основное применение - обучение больших моделей, которые не помещаются в память одного устройства.
Синтаксис
tf.distribute.ParameterServerStrategy(
cluster_resolver=None,
variable_partitioner=None
)
Пример
Давайте создадим экземпляр стратегии с использованием TFConfigClusterResolver
и выведем информацию о количестве реплик:
import tensorflow as tf
resolver = tf.distribute.cluster_resolver.TFConfigClusterResolver()
strategy = tf.distribute.ParameterServerStrategy(resolver)
print(strategy.num_replicas_in_sync)
Результат выполнения кода:
1
Пример
Давайте создадим стратегию и обернем в нее создание переменной модели, а затем выведем значение переменной:
import tensorflow as tf
strategy = tf.distribute.ParameterServerStrategy()
with strategy.scope():
v = tf.Variable([1, 2, 3, 4, 5], name='weights')
print(v)
Результат выполнения кода:
<tf.Variable 'weights:0' shape=(5,) dtype=int32, numpy=array([1, 2, 3, 4, 5], dtype=int32)>
Пример
Давайте обучим простую модель с использованием стратегии и выведем итоговые потери:
import tensorflow as tf
tf.random.set_seed(0)
strategy = tf.distribute.ParameterServerStrategy()
with strategy.scope():
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(1,))
])
model.compile(optimizer='sgd', loss='mse')
x = tf.constant([[1.0], [2.0], [3.0], [4.0]])
y = tf.constant([[2.0], [4.0], [6.0], [8.0]])
history = model.fit(x, y, epochs=5, verbose=0)
print(history.history['loss'][-1])
Результат выполнения кода:
0.02857142873108387
Смотрите также
-
класс
MultiWorkerMirroredStrategy,
который реализует синхронное обучение на нескольких рабочих узлах -
класс
OneDeviceStrategy,
который используется для обучения на одном устройстве -
класс
TPUStrategy,
который предназначен для обучения на TPU -
функцию
get_strategy,
которая возвращает текущую стратегию распределения