tf.contrib.distribute.StandardSingleLossStep
Функция шага, которая реализует шаг обучения для нейронной сети прямого распространения.
Наследуется от: StandardInputStep
tf.contrib.distribute.StandardSingleLossStep(
dataset_fn, loss_fn, optimizer, distribution, iterations_per_step=1
)
Экземпляр этого класса предназначен для использования в качестве вызываемого объекта:
...
step = step_fn.StandardSingleLossStep(
dataset, loss_fn, optimizer, distribution)
# Run a single training step on a given DistributionStrategy:
step(distribution)
...
| Аргументы | |
|---|---|
dataset_fn | функция, которая возвращает tf.data Dataset, который производит входные данные для модели. |
loss_fn | функция, которая принимает контекст и входные данные в качестве аргументов. Она возвращает потерю для этих входных данных. context — это экземпляр values.MultiStepContext, который будет передан при выполнении loss_fn. context может быть использован для указания выходных данных, которые будут возвращены от loss_fn, помимо прочего. |
optimizer | оптимизатор, который реализует правило обновления. |
distribution | объект DistributionStrategy. |
| Атрибуты | |
|---|---|
distribution | |
Методы
initialize
initialize()
__call__
__call__()
Выполнить один шаг этого алгоритма обучения.
© 2020 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 3.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r1.15/api_docs/python/tf/contrib/distribute/StandardSingleLossStep