Функция instance_norm
Функция instance_norm применяет instance-нормализацию к входным данным.
Этот метод нормализует каждый канал каждого объекта в батче независимо,
используя среднее и дисперсию, вычисленные по пространственным измерениям.
Параметр input принимает тензор размерности (N, C, H, W) для
2D данных или (N, C, L) для 1D данных. Параметр running_mean
и running_var используются для накопления статистики в режиме
обучения. Параметр weight и bias задают аффинное преобразование
после нормализации. Параметр momentum управляет обновлением
бегущей статистики. Параметр eps добавляется к дисперсии для
численной стабильности.
Синтаксис
torch.nn.functional.instance_norm(
input,
running_mean=None,
running_var=None,
weight=None,
bias=None,
use_input_stats=True,
momentum=0.1,
eps=1e-05
)
Пример
Базовое применение instance-нормализации к тензору изображений:
import torch
import torch.nn.functional as F
torch.manual_seed(0)
t = torch.randn(2, 3, 2, 2)
res = F.instance_norm(t)
print(res.shape)
Результат выполнения кода:
torch.Size([2, 3, 2, 2])
Пример
Использование с аффинным преобразованием и бегущей статистикой:
import torch
import torch.nn.functional as F
torch.manual_seed(0)
t = torch.randn(2, 2, 3, 3)
weight = torch.ones(2)
bias = torch.zeros(2)
mean = torch.zeros(2)
var = torch.ones(2)
res = F.instance_norm(
t,
running_mean=mean,
running_var=var,
weight=weight,
bias=bias,
use_input_stats=False,
momentum=0.1,
eps=1e-05
)
print(res[0, 0, 0, 0].item())
Результат выполнения кода:
-0.4591662287712097
Пример
Применение instance-нормализации к 1D данным:
import torch
import torch.nn.functional as F
torch.manual_seed(0)
t = torch.randn(3, 4, 5)
res = F.instance_norm(t)
print(res.shape)
print(res[0, 0, 0].item())
Результат выполнения кода:
torch.Size([3, 4, 5])
-0.681307852268219
Пример
Сравнение с ручной реализацией:
import torch
import torch.nn.functional as F
torch.manual_seed(0)
t = torch.randn(2, 3, 2, 2)
res = F.instance_norm(t)
for n in range(2):
for c in range(3):
channel = t[n, c, :, :]
mean = channel.mean()
std = channel.std(unbiased=False)
manual = (channel - mean) / (std + 1e-5)
print(torch.allclose(manual, res[n, c, :, :], atol=1e-4))
Результат выполнения кода:
True
True
True
True
True
True
Смотрите также
-
функцию
batch_norm,
которая нормализует данные по батчу -
функцию
layer_norm,
которая нормализует данные по слою -
функцию
group_norm,
которая разбивает каналы на группы и нормализует их -
функцию
local_response_norm,
которая выполняет локальную нормализацию по каналам