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