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
Предупреждение
Эта функция устарела. Используйте
torch.nn.utils.parametrizations.weight_norm(), которая использует современный API параметризации. Новаяweight_normсовместима сstate_dict, сгенерированными из старыхweight_norm.Руководство по миграции:
- Величина (
weight_g) и направление (weight_v) теперь выражаются какparametrizations.weight.original0иparametrizations.weight.original1соответственно. Если это вас беспокоит, пожалуйста, оставьте комментарий на https://github.com/pytorch/pytorch/issues/102999 - Чтобы удалить перепараметризацию нормализации весов, используйте
torch.nn.utils.parametrize.remove_parametrizations(). - Вес больше не пересчитывается один раз при прямом вызове модуля; вместо этого он будет пересчитываться при каждом доступе. Чтобы восстановить старое поведение, используйте
torch.nn.utils.parametrize.cached()перед вызовом модуля.
- Параметры
- Возвращает
-
Оригинальный модуль с хуком нормализации весов
- Тип возвращаемого значения
-
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/2.1/generated/torch.nn.utils.weight_norm.html