tf.raw_ops.SparseAccumulatorTakeGradient
Извлекает средний разреженный градиент в SparseConditionalAccumulator.
tf.raw_ops.SparseAccumulatorTakeGradient(
handle, num_required, dtype, name=None
)
Операция будет блокироваться до тех пор, пока не будет накоплено достаточное количество (т.е. более чем num_required) градиентов. Если накопитель уже агрегировал более чем num_required градиентов, он вернёт среднее значение накопленных градиентов. Также автоматически увеличивает записанный global_step в накопителе на 1 и сбрасывает агрегированное значение в 0.
| Аргументы | |
|---|---|
handle | A Tensor типа mutable string. Дескриптор SparseConditionalAccumulator. |
num_required | A Tensor типа int32. Количество градиентов, необходимых для возвращения агрегата. |
dtype | tf.DType из: tf.float32, tf.float64, tf.int32, tf.uint8, tf.int16, tf.int8, tf.complex64, tf.int64, tf.qint8, tf.quint8, tf.qint32, tf.bfloat16, tf.uint16, tf.complex128, tf.half, tf.uint32, tf.uint64. Тип данных накопленных градиентов. Должен соответствовать типу накопителя. |
name | Название операции (необязательно). |
| Возвращаемое значение | |
|---|---|
Кортеж из Tensor объектов (индексы, значения, форма). | |
indices | A Tensor типа int64. |
values | A Tensor типа dtype. |
shape | A Tensor типа int64. |
© 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/raw_ops/SparseAccumulatorTakeGradient