РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
316 of 769 menu

Класс 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,
    который нормализует каждый объект в батче независимо
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить