Spec-Zone.ru › PyTorch 2.14

torch.nn.utils.parametrizations.weight_norm

torch.nn.utils.parametrizations.weight_norm(module, name='weight', dim=0) [исходный код]

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

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

Нормализация весов — это репараметризация, которая отделяет величину тензора весов от его направления. Параметр, заданный с помощью name, заменяется двумя параметрами: один задает величину, а другой — направление.

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

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

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

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

Пример:

>>> m = weight_norm(nn.Linear(20, 40), name='weight')
>>> m
ParametrizedLinear(
  in_features=20, out_features=40, bias=True
  (parametrizations): ModuleDict(
    (weight): ParametrizationList(
      (0): _WeightNorm()
    )
  )
)
>>> m.parametrizations.weight.original0.size()
torch.Size([40, 1])
>>> m.parametrizations.weight.original1.size()
torch.Size([40, 20])

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

Spec-Zone.ru

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