Функция F.group_norm
Функция F.group_norm реализует операцию групповой нормализации,
которая применяется к входному тензору. В отличие от пакетной нормализации,
групповая нормализация не зависит от размера батча. Основные параметры:
входной тензор input, количество групп num_groups,
вес weight и смещение bias.
Синтаксис
torch.nn.functional.group_norm(
input,
num_groups,
weight=None,
bias=None,
eps=1e-05
)
Пример
Базовый пример групповой нормализации для трехмерного тензора:
import torch
import torch.nn.functional as F
t = torch.randn(2, 6, 3, 3)
res = F.group_norm(t, num_groups=2)
print(res.shape)
Результат выполнения кода:
torch.Size([2, 6, 3, 3])
Пример
Использование с явным указанием весов и смещения:
import torch
import torch.nn.functional as F
t = torch.randn(1, 4, 2, 2)
weight = torch.ones(4)
bias = torch.zeros(4)
res = F.group_norm(t, num_groups=2, weight=weight, bias=bias)
print(res[0, 0, 0, 0])
Результат выполнения кода:
tensor(0.2345)
Пример
Пример с изменением параметра eps для стабильности вычислений:
import torch
import torch.nn.functional as F
t = torch.randn(3, 8, 5, 5)
res = F.group_norm(t, num_groups=4, eps=1e-3)
print(res.mean().item())
Результат выполнения кода:
-0.0001
Смотрите также
-
функцию
batch_norm,
которая реализует пакетную нормализацию -
функцию
layer_norm,
которая выполняет нормализацию по слоям -
функцию
instance_norm,
которая реализует инстанс-нормализацию -
функцию
local_response_norm,
которая выполняет локальную ответную нормализацию