Метод distribute_datasets_from_function
Метод distribute_datasets_from_function класса MirroredStrategy
создает распределенный датасет из пользовательской функции.
Первый параметр - это функция, которая принимает объект
InputContext и возвращает датасет tf.data.Dataset.
Второй параметр - необязательные опции распределения.
Такой подход позволяет гибко управлять разбиением данных
между устройствами и отличается от метода
experimental_distribute_dataset,
который распределяет уже готовый датасет.
Синтаксис
strategy.distribute_datasets_from_function(
dataset_fn,
options=None
)
Пример
Давайте создадим стратегию MirroredStrategy
и распределим датасет, используя функцию,
которая принимает контекст и возвращает датасет из чисел:
import tensorflow as tf
tf.random.set_seed(0)
strategy = tf.distribute.MirroredStrategy()
def dataset_fn(input_context):
batch_size = input_context.get_per_replica_batch_size(4)
dataset = tf.data.Dataset.from_tensor_slices([1, 2, 3, 4, 5])
dataset = dataset.batch(batch_size)
return dataset
distributed_dataset = strategy.distribute_datasets_from_function(dataset_fn)
for batch in distributed_dataset:
print(batch)
Результат выполнения кода:
tf.Tensor([1 2 3 4 5], shape=(5,), dtype=int32)
Пример
Давайте распределим датасет внутри области видимости стратегии и выполним обучение простой модели на полученных данных:
import tensorflow as tf
tf.random.set_seed(0)
strategy = tf.distribute.MirroredStrategy()
def dataset_fn(input_context):
batch_size = input_context.get_per_replica_batch_size(2)
dataset = tf.data.Dataset.from_tensor_slices(
[[1, 2, 3], [4, 5, 6]]
)
dataset = dataset.batch(batch_size)
return dataset
with strategy.scope():
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(3,))
])
model.compile(optimizer='sgd', loss='mse')
distributed_dataset = strategy.distribute_datasets_from_function(dataset_fn)
for batch in distributed_dataset:
print(batch)
Результат выполнения кода:
tf.Tensor(
[[1 2 3]
[4 5 6]], shape=(2, 3), dtype=int32)
Смотрите также
-
класс
MirroredStrategy,
который реализует зеркальное распределение -
метод
scope,
который создает контекст для распределенных переменных -
метод
run,
который выполняет функцию на каждом устройстве -
метод
experimental_distribute_dataset,
который распределяет готовый датасет между устройствами