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)
© 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/r2.3/api_docs/python/tf/clip_by_global_norm