Spec-Zone.ru › PyTorch 1

torch.nn.utils.weight_norm

torch.nn.utils.weight_norm(module, name='weight', dim=0) [source]

Применяет нормализацию весов к параметру в данном модуле.

w=gv∥v∥\mathbf{w} = g \dfrac{\mathbf{v}}{\|\mathbf{v}\|}

Нормализация весов — это перепараметризация, которая отделяет величину тензора весов от его направления. Это заменяет параметр, указанный name (например, 'weight') на два параметра: один, определяющий величину (например, 'weight_g'), и один, определяющий направление (например, 'weight_v'). Нормализация весов реализуется с помощью хука, который пересчитывает тензор весов из величины и направления перед каждым forward() вызовом.

По умолчанию, с dim=0, норма вычисляется независимо для каждого выходного канала/плоскости. Для вычисления нормы по всему тензору весов используйте dim=None.

См. https://arxiv.org/abs/1602.07868

Параметры:
  • module (Модуль) – содержащий модуль
  • name (str, необязательно) – имя параметра веса
  • dim (int, необязательно) – размерность, по которой вычислять норму
Возвращает:

Исходный модуль с хуком нормализации весов

Тип возвращаемого значения:

T_module

Пример:

>>> m = weight_norm(nn.Linear(20, 40), name='weight')
>>> m
Linear(in_features=20, out_features=40, bias=True)
>>> m.weight_g.size()
torch.Size([40, 1])
>>> m.weight_v.size()
torch.Size([40, 20])

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.nn.utils.weight_norm.html

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API