tf.compat.v1.ConditionalAccumulator
Условный аккумулирующий градиенты.
Наследуется от: ConditionalAccumulatorBase
tf.compat.v1.ConditionalAccumulator(
dtype,
shape=None,
shared_name=None,
name='conditional_accumulator',
reduction_type='MEAN'
)
Актуальные градиенты (т.е. шаг времени, на котором был вычислен градиент, равен шагу времени аккумулирующего устройства) добавляются в аккумулирующее устройство.
Извлечение среднего градиента блокируется до тех пор, пока не будет накоплено необходимое количество градиентов.
| Аргументы | |
|---|---|
dtype | Тип данных аккумулируемых градиентов. |
shape | Форма аккумулируемых градиентов. |
shared_name | Необязательно. Если не пусто, этот аккумулирующий объект будет совместно использоваться под заданным именем в нескольких сессиях. |
name | Необязательное имя для аккумулирующего устройства. |
reduction_type | Тип редукции, используемый при получении градиента. |
| Атрибуты | |
|---|---|
accumulator_ref | Ссылка на базовое аккумулирующее устройство. |
dtype | Тип данных градиентов, накапливаемых этим аккумулирующим устройством. |
name | Имя базового аккумулирующего устройства. |
Методы
apply_grad
apply_grad(
grad, local_step=0, name=None
)
Попытка применить градиент к аккумулирующему устройству.
Попытка будет проигнорирована, если градиент устарел, т.е. локальный шаг меньше глобального шага времени аккумулирующего устройства.
| Аргументы | |
|---|---|
grad | Градиент-тензор, который нужно применить. |
local_step | Шаг времени, на котором был вычислен градиент. |
name | Необязательное имя операции. |
| Возвращаемые значения | |
|---|---|
| Операция, которая (условно) применяет градиент к аккумулирующему устройству. |
| Исключения | |
|---|---|
ValueError | Если форма grad неверна |
num_accumulated
num_accumulated(
name=None
)
Количество градиентов, которые в настоящее время были агрегированы в аккумулирующем устройстве.
| Аргументы | |
|---|---|
name | Необязательное имя операции. |
| Возвращаемые значения | |
|---|---|
| Количество накопленных градиентов в настоящее время в аккумулирующем устройстве. |
set_global_step
set_global_step(
new_global_step, name=None
)
Устанавливает глобальный шаг времени аккумулирующего устройства.
Операция выводит предупреждение, если мы пытаемся установить шаг времени, который меньше, чем собственный шаг времени аккумулирующего устройства.
| Аргументы | |
|---|---|
new_global_step | Значение нового шага времени. Может быть переменной или константой |
name | Необязательное имя операции. |
| Возвращаемые значения | |
|---|---|
| Операция, которая устанавливает шаг времени аккумулирующего устройства. |
take_grad
take_grad(
num_required, name=None
)
Попытка извлечь средний градиент из аккумулирующего устройства.
Операция блокируется до тех пор, пока достаточное количество градиентов не будет успешно применено к аккумулирующему устройству.
После успешного выполнения выполняются также следующие действия:
- Счётчик накопленных градиентов сбрасывается до 0.
- Агрегированный градиент сбрасывается до 0 тензора.
- Внутренний шаг времени аккумулирующего устройства увеличивается на 1.
| Аргументы | |
|---|---|
num_required | Количество градиентов, которые должны быть агрегированы |
name | Необязательное имя операции |
| Возвращаемые значения | |
|---|---|
| Тензор, содержащий значение среднего градиента. |
| Исключения | |
|---|---|
InvalidArgumentError | Если num_required < 1 |
© 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/api_docs/python/tf/compat/v1/ConditionalAccumulator