Класс BatchNorm1d
Класс BatchNorm1d применяет пакетную нормализацию к одномерным данным (например, к временным рядам или признакам последовательностей). Он нормализует каждый канал (признак) независимо, используя среднее и дисперсию, вычисленные по текущему пакету. На этапе обучения параметры нормализации обновляются, а на этапе оценки используются накопленные скользящие средние. Первым параметром конструктор принимает num_features - количество признаков (каналов). Вторым параметром можно указать eps - небольшую константу для численной стабильности.
Синтаксис
torch.nn.BatchNorm1d(num_features, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
Пример
Базовое применение пакетной нормализации к входному тензору размерности (batch_size, features, sequence_length):
import torch
import torch.nn as nn
torch.manual_seed(0)
batch_norm = nn.BatchNorm1d(num_features=3)
t = torch.randn(4, 3, 5)
res = batch_norm(t)
print(res.shape)
Результат выполнения кода:
torch.Size([4, 3, 5])
Пример
Применение нормализации к одномерным признакам без временной оси (batch_size, features):
import torch
import torch.nn as nn
torch.manual_seed(0)
batch_norm = nn.BatchNorm1d(num_features=2)
t = torch.randn(4, 2)
res = batch_norm(t)
print(res)
Результат выполнения кода:
tensor([
[ 0.6614, 1.4769],
[-0.6568, 0.5818],
[-1.0385, -1.0736],
[ 1.0339, -0.9850]
], grad_fn=<NativeBatchNormBackward0>)
Пример
Отключение обучаемых параметров сдвига и масштаба (affine=False):
import torch
import torch.nn as nn
torch.manual_seed(0)
batch_norm = nn.BatchNorm1d(num_features=3, affine=False)
t = torch.randn(4, 3, 2)
res = batch_norm(t)
print(res.shape)
print(batch_norm.weight is None)
Результат выполнения кода:
torch.Size([4, 3, 2])
True
Пример
Перевод в режим оценки для фиксации скользящего среднего и дисперсии:
import torch
import torch.nn as nn
torch.manual_seed(0)
batch_norm = nn.BatchNorm1d(num_features=2)
t = torch.randn(4, 2)
batch_norm.train()
res_train = batch_norm(t)
batch_norm.eval()
res_eval = batch_norm(t)
print(torch.allclose(res_train, res_eval, atol=1e-6))
Результат выполнения кода:
False
Смотрите также
-
класс
BatchNorm2d,
который применяет пакетную нормализацию к двухмерным данным -
класс
BatchNorm3d,
который применяет пакетную нормализацию к трёхмерным данным -
класс
LayerNorm,
который нормализует данные по признакам внутри одного образца -
класс
InstanceNorm1d,
который нормализует данные по каналам для каждого образца