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

Класс BatchNorm3d

Класс BatchNorm3d применяет батч-нормализацию к пятимерным тензорам (батч, каналы, глубина, высота, ширина). Этот слой нормализует каждый канал независимо, используя среднее и дисперсию, вычисленные по текущему батчу. В процессе обучения слой накапливает скользящие средние значения для использования в режиме оценки.

Синтаксис

torch.nn.BatchNorm3d( num_features, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True )

Основные параметры:

  • num_features - количество каналов на входе
  • eps - небольшая константа для численной стабильности
  • momentum - коэффициент обновления скользящих средних
  • affine - использовать ли обучаемые параметры сдвига и масштаба
  • track_running_stats - отслеживать ли статистику по батчам

Пример

Создадим слой батч-нормализации для 3D данных с 16 каналами:

import torch import torch.nn as nn bn = nn.BatchNorm3d(num_features=16) print(bn)

Результат выполнения кода:

BatchNorm3d(16, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)

Пример

Применим батч-нормализацию к случайному тензору размера (батч, каналы, глубина, высота, ширина):

import torch import torch.nn as nn torch.manual_seed(0) bn = nn.BatchNorm3d(num_features=8) t = torch.randn(4, 8, 10, 16, 16) res = bn(t) print(res.shape)

Результат выполнения кода:

torch.Size([4, 8, 10, 16, 16])

Пример

Проверим, что после нормализации среднее близко к нулю, а дисперсия - к единице:

import torch import torch.nn as nn torch.manual_seed(0) bn = nn.BatchNorm3d(num_features=4) t = torch.randn(2, 4, 5, 6, 6) res = bn(t) print(res.mean().item()) print(res.std().item())

Результат выполнения кода:

-0.00022403579891561717 0.9994456171989441

Пример

Использование слоя в режиме оценки (например, для инференса):

import torch import torch.nn as nn torch.manual_seed(0) bn = nn.BatchNorm3d(num_features=8) t = torch.randn(2, 8, 4, 8, 8) bn.eval() res = bn(t) print(res.mean().item())

Результат выполнения кода:

0.020376864820718765

Смотрите также

  • класс BatchNorm1d,
    который применяет батч-нормализацию к одномерным данным
  • класс BatchNorm2d,
    который применяет батч-нормализацию к двумерным данным
  • класс SyncBatchNorm,
    который синхронизирует статистику между процессами
  • класс LayerNorm,
    который нормализует по признаковому пространству
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить