Класс OneDeviceStrategy
Класс OneDeviceStrategy создает стратегию распределения,
которая размещает все переменные и вычисления на одном
устройстве. Это удобно для отладки и прототипирования кода,
который позже будет выполняться на нескольких устройствах,
поскольку позволяет писать код единообразно. Первым параметром
конструктор принимает строку с именем устройства, например
'/cpu:0' или '/gpu:0'.
Синтаксис
tf.distribute.OneDeviceStrategy(device)
Пример
Давайте создадим стратегию для процессора и выполним в ее контексте простое умножение тензора:
import tensorflow as tf
strategy = tf.distribute.OneDeviceStrategy(device='/cpu:0')
with strategy.scope():
t = tf.constant([1, 2, 3, 4, 5])
res = t * 2
print(res)
Результат выполнения кода:
tf.Tensor([ 2 4 6 8 10], shape=(5,), dtype=int32)
Пример
Давайте создадим модель внутри области видимости стратегии и проверим, что переменные размещаются на нужном устройстве:
import tensorflow as tf
strategy = tf.distribute.OneDeviceStrategy(device='/cpu:0')
with strategy.scope():
model = tf.keras.Sequential([
tf.keras.layers.Dense(3, input_shape=(2,))
])
model.compile(optimizer='sgd', loss='mse')
print(strategy.num_replicas_in_sync)
print(model.weights[0].device)
Результат выполнения кода:
1
"/job:localhost/replica:0/task:0/device:CPU:0"
Пример
Давайте выполним обучение небольшой модели в контексте стратегии одного устройства:
import tensorflow as tf
tf.random.set_seed(0)
strategy = tf.distribute.OneDeviceStrategy(device='/cpu:0')
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.0123456789
Смотрите также
-
функцию
get_strategy,
которая возвращает текущую стратегию распределения -
функцию
has_strategy,
которая проверяет наличие активной стратегии -
класс
TPUStrategy,
который создает стратегию для работы с TPU -
функцию
config.list_logical_devices,
которая возвращает список логических устройств