Функция F.batch_norm
Функция F.batch_norm применяет пакетную нормализацию (batch normalization) к входному тензору. Она нормализует каждый канал входных данных по среднему и стандартному отклонению, вычисленным по пакету. Первым параметром функция принимает входной тензор input, вторым и третьим - тензоры среднего running_mean и дисперсии running_var, накопленные во время обучения. Четвертым параметром передается тензор весов weight, пятым - смещение bias. Шестым параметром можно задать коэффициент момента momentum, седьмым - значение эпсилон eps для численной стабильности, а восьмым - флаг training, определяющий режим работы (обучение или оценка).
Синтаксис
torch.nn.functional.batch_norm(
input,
running_mean,
running_var,
weight=None,
bias=None,
training=False,
momentum=0.1,
eps=1e-05
)
Пример
Выполним пакетную нормализацию входного тензора размером (2, 3, 4, 4) в режиме оценки:
import torch
import torch.nn.functional as F
t = torch.randn(2, 3, 4, 4)
running_mean = torch.zeros(3)
running_var = torch.ones(3)
weight = torch.ones(3)
bias = torch.zeros(3)
res = F.batch_norm(
t,
running_mean,
running_var,
weight,
bias,
training=False
)
print(res.shape)
Результат выполнения кода:
torch.Size([2, 3, 4, 4])
Пример
Используем пакетную нормализацию в режиме обучения с обновлением накопленных статистик:
import torch
import torch.nn.functional as F
torch.manual_seed(0)
t = torch.randn(4, 2, 3, 3)
running_mean = torch.zeros(2)
running_var = torch.ones(2)
weight = torch.ones(2)
bias = torch.zeros(2)
res = F.batch_norm(
t,
running_mean,
running_var,
weight,
bias,
training=True,
momentum=0.1,
eps=1e-05
)
print(res.shape)
print(running_mean)
print(running_var)
Результат выполнения кода:
torch.Size([4, 2, 3, 3])
tensor([0.0392, 0.0181])
tensor([1.0588, 1.0140])
Пример
Применяем пакетную нормализацию к одномерному входному тензору:
import torch
import torch.nn.functional as F
t = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0])
running_mean = torch.tensor([0.0])
running_var = torch.tensor([1.0])
weight = torch.tensor([2.0])
bias = torch.tensor([1.0])
res = F.batch_norm(
t,
running_mean,
running_var,
weight,
bias,
training=False
)
print(res)
Результат выполнения кода:
tensor([3.0000, 5.0000, 7.0000, 9.0000, 11.0000])
Смотрите также
-
функцию
F.layer_norm,
которая применяет слоевую нормализацию -
функцию
F.group_norm,
которая применяет групповую нормализацию -
функцию
F.instance_norm,
которая применяет нормализацию по экземплярам -
функцию
F.local_response_norm,
которая применяет локальную нормализацию по ответам