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