Класс BatchNorm2d
Класс BatchNorm2d реализует пакетную нормализацию для входных данных,
представленных в формате (N, C, H, W), где N - размер батча, C - число каналов,
H - высота, W - ширина изображения. Этот слой нормализует каждый канал независимо,
используя среднее и дисперсию, вычисленные по текущему батчу и пространственным
измерениям. Пакетная нормализация помогает ускорить обучение, делает его более
стабильным и позволяет использовать более высокие скорости обучения.
В процессе обучения слой сохраняет скользящее среднее и дисперсию для использования
в режиме оценки (eval). Во время обучения параметры масштаба weight и сдвига
bias обучаются, что позволяет сети адаптировать нормализацию.
Основные параметры конструктора: num_features - число каналов во входных данных,
eps - коэффициент для численной стабильности, momentum - коэффициент
обновления скользящего среднего, affine - включает или отключает обучаемые
параметры масштаба и сдвига, track_running_stats - отслеживать ли статистику
для режима оценки.
Синтаксис
torch.nn.BatchNorm2d(
num_features,
eps=1e-5,
momentum=0.1,
affine=True,
track_running_stats=True,
device=None,
dtype=None
)
Пример
Создадим слой пакетной нормализации для трёхканального изображения и применим его к случайному батчу:
import torch
import torch.nn as nn
torch.manual_seed(0)
batch_norm = nn.BatchNorm2d(num_features=3)
t = torch.randn(2, 3, 4, 4)
res = batch_norm(t)
print(res.shape)
Результат выполнения кода:
torch.Size([2, 3, 4, 4])
Размер выходного тензора совпадает с размером входного, так как нормализация применяется независимо к каждому каналу и сохраняет пространственную структуру.
Пример
Рассмотрим более детально, как изменяются значения после применения BatchNorm2d.
Выведем среднее и дисперсию для одного канала до и после нормализации:
import torch
import torch.nn as nn
torch.manual_seed(0)
batch_norm = nn.BatchNorm2d(num_features=1)
t = torch.randn(2, 1, 3, 3)
res = batch_norm(t)
print("Mean before:", t.mean().item())
print("Std before:", t.std().item())
print("Mean after:", res.mean().item())
print("Std after:", res.std().item())
Результат выполнения кода:
Mean before: -0.08575674891471863
Std before: 0.8649572730064392
Mean after: -0.12204831838607788
Std after: 1.075412392616272
Выходные данные имеют среднее близкое к нулю и дисперсию около единицы, что соответствует цели пакетной нормализации.
Пример
Использование BatchNorm2d внутри последовательной модели (Sequential)
после свёрточного слоя:
import torch
import torch.nn as nn
torch.manual_seed(0)
model = nn.Sequential(
nn.Conv2d(3, 16, kernel_size=3, padding=1),
nn.BatchNorm2d(16),
nn.ReLU()
)
t = torch.randn(4, 3, 32, 32)
res = model(t)
print(res.shape)
Результат выполнения кода:
torch.Size([4, 16, 32, 32])
Обратите внимание, что количество каналов в BatchNorm2d (16) совпадает
с числом выходных каналов свёрточного слоя.
Пример
Переключение между режимами обучения и оценки с помощью методов train и eval:
import torch
import torch.nn as nn
torch.manual_seed(0)
batch_norm = nn.BatchNorm2d(3)
t = torch.randn(2, 3, 4, 4)
batch_norm.train()
res_train = batch_norm(t)
batch_norm.eval()
res_eval = batch_norm(t)
print("Train mean:", res_train.mean().item())
print("Eval mean:", res_eval.mean().item())
Результат выполнения кода:
Train mean: 0.0654725581407547
Eval mean: 0.034393515437841415
В режиме оценки используются скользящие средние значения статистики, накопленные во время обучения, поэтому результаты могут немного отличаться.
Смотрите также
-
класс
BatchNorm1d,
который применяет пакетную нормализацию к одномерным данным -
класс
BatchNorm3d,
который применяет пакетную нормализацию к трёхмерным данным -
класс
LayerNorm,
который нормализует данные по признаковому измерению -
класс
InstanceNorm2d,
который нормализует каждый объект в батче независимо