Класс InstanceNorm1d
Класс InstanceNorm1d применяет нормализацию по каналам для каждого отдельного объекта в мини-батче.
В отличие от BatchNorm1d, которая использует статистику по батчу, этот метод вычисляет среднее и дисперсию для каждого объекта и каждого канала отдельно.
Параметры: num_features - количество признаков (каналов) во входных данных; eps - значение для численной стабильности (по умолчанию 1e-5);
momentum - коэффициент для скользящего среднего (используется, если track_running_stats=True);
affine - флаг, определяющий, нужно ли обучать параметры сдвига и масштаба (True по умолчанию);
track_running_stats - флаг, отслеживать ли статистику для фазы оценки (True по умолчанию).
Синтаксис
torch.nn.InstanceNorm1d(num_features, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
Пример
Базовое применение InstanceNorm1d к тензору формы (N, C, L):
import torch
import torch.nn as nn
t = torch.tensor([
[[1.0, 2.0, 3.0]],
[[4.0, 5.0, 6.0]]
])
norm = nn.InstanceNorm1d(1)
res = norm(t)
print(res)
Результат выполнения кода:
tensor([
[[-1.2247, 0.0000, 1.2247]],
[[-1.2247, 0.0000, 1.2247]]
], grad_fn=<NativeInstanceNormBackward0>)
Пример
Использование с отключенным параметром affine (без обучаемых весов):
import torch
import torch.nn as nn
t = torch.tensor([
[[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]],
[[7.0, 8.0, 9.0], [10.0, 11.0, 12.0]]
])
norm = nn.InstanceNorm1d(2, affine=False)
res = norm(t.float())
print(res)
Результат выполнения кода:
tensor([
[[-1.2247, 0.0000, 1.2247],
[-1.2247, 0.0000, 1.2247]],
[[-1.2247, 0.0000, 1.2247],
[-1.2247, 0.0000, 1.2247]]
])
Пример
Применение слоя в составе последовательной модели:
import torch
import torch.nn as nn
model = nn.Sequential(
nn.Linear(10, 5),
nn.InstanceNorm1d(5),
nn.ReLU()
)
t = torch.randn(3, 10)
res = model(t)
print(res.shape)
Результат выполнения кода:
torch.Size([3, 5])
Смотрите также
-
класс
BatchNorm1d,
который нормализует данные по батчу -
класс
InstanceNorm2d,
который применяет Instance Normalization к двумерным данным -
класс
LayerNorm,
который нормализует данные по признакам -
класс
GroupNorm,
который разделяет каналы на группы для нормализации