Spec-Zone.ru › TensorFlow 2.9

tf.keras.utils.experimental.DatasetCreator

Объект, возвращающий tf.data.Dataset при вызове.

tf.keras.utils.experimental.DatasetCreator(
    dataset_fn, input_options=None
)

tf.keras.utils.experimental.DatasetCreator обозначен как поддерживаемый тип для x, или входных данных, в tf.keras.Model.fit. Передайте экземпляр этого класса в fit при использовании вызываемого объекта (с аргументом input_context), возвращающего tf.data.Dataset.

model = tf.keras.Sequential([tf.keras.layers.Dense(10)])
model.compile(tf.keras.optimizers.SGD(), loss="mse")

def dataset_fn(input_context):
  global_batch_size = 64
  batch_size = input_context.get_per_replica_batch_size(global_batch_size)
  dataset = tf.data.Dataset.from_tensors(([1.], [1.])).repeat()
  dataset = dataset.shard(
      input_context.num_input_pipelines, input_context.input_pipeline_id)
  dataset = dataset.batch(batch_size)
  dataset = dataset.prefetch(2)
  return dataset

input_options = tf.distribute.InputOptions(
    experimental_fetch_to_device=True,
    experimental_per_replica_buffer_size=2)
model.fit(tf.keras.utils.experimental.DatasetCreator(
    dataset_fn, input_options=input_options), epochs=10, steps_per_epoch=10)

Model.fit с использованием DatasetCreator предназначен для работы со всеми tf.distribute.Strategy, если Strategy.scope используется при создании модели:

strategy = tf.distribute.experimental.ParameterServerStrategy(
    cluster_resolver)
with strategy.scope():
  model = tf.keras.Sequential([tf.keras.layers.Dense(10)])
model.compile(tf.keras.optimizers.SGD(), loss="mse")

def dataset_fn(input_context):
  ...

input_options = ...
model.fit(tf.keras.utils.experimental.DatasetCreator(
    dataset_fn, input_options=input_options), epochs=10, steps_per_epoch=10)
Примечание: При использовании DatasetCreator, аргумент steps_per_epoch в Model.fit должен быть предоставлен, так как кардинальность такого входа не может быть определена.
Аргументы
dataset_fn Функция, принимающая один аргумент типа tf.distribute.InputContext, который используется для вычисления размера пакета и фрагментации входной конвейера между рабочими узлами (если ни то, ни другое не требуется, параметр InputContext можно пропустить в dataset_fn), и возвращающая tf.data.Dataset.
input_options Необязательные tf.distribute.InputOptions, используемые для специфических опций при использовании с распределением, например, для предварительной выборки элементов набора данных в оперативную память ускорителя или оперативную память хоста, а также для размера буфера предварительной выборки в оперативной памяти устройства реплики. Не оказывает никакого влияния, если не используется с распределённым обучением. См. tf.distribute.InputOptions для получения дополнительной информации.

Методы

__call__

Посмотреть исходный код

__call__(
    *args, **kwargs
)

Вызвать self как функцию.

© 2022 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 4.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r2.9/api_docs/python/tf/keras/utils/experimental/DatasetCreator

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API