Класс GroupNorm
Класс GroupNorm применяет групповую нормализацию к входным данным.
Основная идея метода заключается в разделении каналов на группы и вычислении
среднего и дисперсии внутри каждой группы независимо.
Это позволяет нормализовать активации в пределах группы каналов, что особенно
полезно для моделей с ограниченным размером батча.
Первым параметром передаётся количество групп num_groups,
вторым - количество каналов num_channels.
Синтаксис
torch.nn.GroupNorm(num_groups, num_channels, eps=1e-5, affine=True)
Параметры:
num_groups - количество групп, на которые делятся каналы (должно делиться на число каналов);
num_channels - количество каналов во входных данных;
eps - значение для численной стабильности (по умолчанию 1e-5);
affine - булевый параметр, указывающий, следует ли использовать обучаемые параметры
сдвига и масштаба (по умолчанию True).
Пример
Создадим слой групповой нормализации с двумя группами и применим его к тензору:
import torch
import torch.nn as nn
# Создаем слой GroupNorm: 2 группы, 4 канала
gn = nn.GroupNorm(num_groups=2, num_channels=4)
# Входной тензор: батч размером 1, 4 канала, 3x3
t = torch.randn(1, 4, 3, 3)
res = gn(t)
print(res.shape)
Результат выполнения кода:
torch.Size([1, 4, 3, 3])
Пример
Применим групповую нормализацию к трёхмерному тензору с явным указанием параметров:
import torch
import torch.nn as nn
torch.manual_seed(0)
# Создаем слой GroupNorm с 3 группами и 6 каналами
gn = nn.GroupNorm(num_groups=3, num_channels=6, eps=1e-6, affine=True)
# Входной тензор размером 2x6x4x4
t = torch.randn(2, 6, 4, 4)
res = gn(t)
print(res[0, 0, 0, :5])
Результат выполнения кода:
tensor([-0.2377, 0.8390, -1.5702, 1.1739, -0.1737])
Пример
Рассмотрим, как GroupNorm работает с одномерными данными:
import torch
import torch.nn as nn
# Слой GroupNorm для одномерных данных
gn = nn.GroupNorm(num_groups=2, num_channels=4)
# Входной тензор: батч 3, 4 канала, 10 элементов
t = torch.randn(3, 4, 10)
res = gn(t)
print(res.shape)
print(t.mean(dim=(2,), keepdim=True).shape)
Результат выполнения кода:
torch.Size([3, 4, 10])
torch.Size([3, 4, 1])
Смотрите также
-
класс
BatchNorm2d,
который применяет пакетную нормализацию к данным -
класс
LayerNorm,
который применяет нормализацию по признаковому пространству -
класс
InstanceNorm2d,
который применяет нормализацию к каждому экземпляру отдельно -
класс
SyncBatchNorm,
который синхронизирует пакетную нормализацию между устройствами