РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
814 of 824 menu

Класс 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,
    которая возвращает текущую стратегию распределения
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить