torch.nn.utils.parametrizations.weight_norm
-
torch.nn.utils.parametrizations.weight_norm(module, name='weight', dim=0)[исходный код] -
Применяет нормализацию весов к параметру указанного модуля.
Нормализация весов — это репараметризация, которая отделяет величину тензора весов от его направления. Параметр, заданный с помощью
name, заменяется двумя параметрами: один задает величину, а другой — направление.По умолчанию, при использовании
dim=0норма вычисляется независимо для каждого выходного канала/плоскости. Чтобы вычислить норму по всему тензору весов, используйтеdim=None.См. https://arxiv.org/abs/1602.07868
- Параметры:
- Возвращает:
-
Исходный модуль с хуком нормализации весов
Пример:
>>> 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