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

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