Функция has_strategy
Функция has_strategy проверяет, определена ли в текущем контексте выполнения стратегия распределенного обучения. Функция не принимает параметров и возвращает булево значение: True, если стратегия установлена, и False в противном случае. Обычно она используется внутри функций, которые должны по-разному работать в режиме распределенного и обычного обучения.
Синтаксис
tf.distribute.has_strategy()
Пример
Давайте проверим наличие стратегии вне контекста распределенного обучения:
import tensorflow as tf
res = tf.distribute.has_strategy()
print(res)
Результат выполнения кода:
False
Пример
Давайте проверим наличие стратегии внутри контекста OneDeviceStrategy:
import tensorflow as tf
strategy = tf.distribute.OneDeviceStrategy("/cpu:0")
with strategy.scope():
res = tf.distribute.has_strategy()
print(res)
Результат выполнения кода:
True
Пример
Давайте используем функцию внутри модели, чтобы выбрать способ вычисления в зависимости от наличия стратегии:
import tensorflow as tf
def compute_loss(y_true, y_pred):
if tf.distribute.has_strategy():
return tf.reduce_mean(tf.square(y_true - y_pred))
return tf.reduce_sum(tf.square(y_true - y_pred))
strategy = tf.distribute.OneDeviceStrategy("/cpu:0")
with strategy.scope():
loss = compute_loss(tf.constant([1.0, 2.0]), tf.constant([1.5, 2.5]))
print(loss)
Результат выполнения кода:
tf.Tensor(0.5, shape=(), dtype=float32)
Смотрите также
-
функцию
get_strategy,
которая возвращает текущую стратегию распределенного обучения -
класс
OneDeviceStrategy,
который задает стратегию обучения на одном устройстве -
класс
MultiWorkerMirroredStrategy,
который задает стратегию обучения на нескольких рабочих узлах -
функцию
device,
которая задает устройство для выполнения операций