tf.clip_by_global_norm
| Просмотреть исходный код на GitHub |
Ограничивает значения нескольких тензоров по отношению суммы их норм.
tf.clip_by_global_norm(
t_list, clip_norm, use_norm=None, name=None
)
Для заданной группы или списка тензоров t_list, и коэффициента ограничения clip_norm, эта операция возвращает список ограниченных тензоров list_clipped и глобальную норму (global_norm) всех тензоров в t_list. Необязательно, если вы уже вычислили глобальную норму для t_list, вы можете указать глобальную норму с помощью use_norm.
Для выполнения ограничения значения t_list[i] устанавливаются равными:
t_list[i] * clip_norm / max(global_norm, clip_norm)
где:
global_norm = sqrt(sum([l2norm(t)**2 for t in t_list]))
Если clip_norm > global_norm, то элементы в t_list остаются неизменными, в противном случае все они уменьшаются на глобальный коэффициент.
Если global_norm == infinity, то элементы в t_list устанавливаются в NaN, чтобы сигнализировать об ошибке.
Любые элементы в t_list, которые имеют тип None, игнорируются.
Это правильный способ выполнения обрезки градиентов (Pascanu et al., 2012).
Однако, он медленнее, чем clip_by_norm(), потому что все параметры должны быть готовы перед выполнением операции обрезки.
| Аргументы | |
|---|---|
t_list | Кортеж или список смешанных Tensors, IndexedSlices, или None. |
clip_norm | 0-мерный (скалярный) Tensor > 0. Коэффициент ограничения. |
use_norm | 0-мерный (скалярный) Tensor типа float (необязательно). Глобальная норма для использования. Если не указано, используется global_norm() для вычисления нормы. |
name | Имя операции (необязательно). |
| Возвращаемые значения | |
|---|---|
list_clipped | Список Tensors того же типа, что и list_t. |
global_norm | 0-мерный (скалярный) Tensor, представляющий глобальную норму. |
| Исключения | |
|---|---|
TypeError | Если t_list не является последовательностью. |
Ссылки:
О трудностях обучения рекуррентным нейронным сетям: Pascanu et al., 2012 (pdf)
© 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/clip_by_global_norm