Класс RMSNorm
Класс RMSNorm применяет нормализацию по среднеквадратичному отклонению к входным данным.
Этот метод нормализует каждый элемент тензора по последней размерности,
используя среднеквадратичное значение вместо дисперсии, что делает его более эффективным,
чем стандартная нормализация слоя. Основное отличие от LayerNorm заключается
в том, что RMSNorm не вычитает среднее значение,
а только масштабирует входные данные по среднеквадратичному отклонению.
Класс принимает размер нормализуемой оси и опционально обучаемый параметр масштаба.
Синтаксис
torch.nn.RMSNorm(normalized_shape, eps=1e-5, elementwise_affine=True, device=None, dtype=None)
Пример
Давайте создадим слой RMSNorm и применим его к простому тензору:
import torch
rms_norm = torch.nn.RMSNorm(3)
t = torch.tensor([[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]])
res = rms_norm(t)
print(res)
Результат выполнения кода:
tensor([
[-1.2247, 0.0000, 1.2247],
[-1.2247, 0.0000, 1.2247],
])
Пример
Теперь создадим слой без обучаемого параметра масштаба, установив elementwise_affine в False:
import torch
rms_norm = torch.nn.RMSNorm(3, elementwise_affine=False)
t = torch.tensor([[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]])
res = rms_norm(t)
print(res)
Результат выполнения кода:
tensor([
[-1.2247, 0.0000, 1.2247],
[-1.2247, 0.0000, 1.2247],
])
Пример
Попробуем изменить параметр eps, который добавляется к знаменателю для численной стабильности:
import torch
rms_norm = torch.nn.RMSNorm(3, eps=1e-3)
t = torch.tensor([[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]])
res = rms_norm(t)
print(res)
Результат выполнения кода:
tensor([
[-1.2194, 0.0000, 1.2194],
[-1.2194, 0.0000, 1.2194],
])
Пример
Используем RMSNorm внутри последовательной модели нейронной сети:
import torch
model = torch.nn.Sequential(
torch.nn.Linear(10, 5),
torch.nn.RMSNorm(5),
torch.nn.ReLU(),
)
t = torch.randn(3, 10)
res = model(t)
print(res.shape)
Результат выполнения кода:
torch.Size([3, 5])
Смотрите также
-
класс
LayerNorm,
который выполняет стандартную нормализацию слоя с вычитанием среднего -
класс
BatchNorm1d,
который применяет пакетную нормализацию для одномерных данных -
класс
GroupNorm,
который группирует каналы для нормализации -
класс
InstanceNorm1d,
который выполняет нормализацию для каждого экземпляра отдельно