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

Класс SyncBatchNorm

Класс SyncBatchNorm из модуля torch.nn предназначен для пакетной нормализации в распределённых сценариях. В отличие от BatchNorm, он синхронизирует среднее и дисперсию по всем процессам (GPU) во время обучения, что позволяет корректно вычислять статистики при малом размере батча на каждом устройстве. Класс принимает такие же параметры, как и BatchNorm: количество признаков, обучаемые параметры affine, момент для вычисления скользящего среднего momentum и коэффициент для численной стабильности eps.

Синтаксис

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

Пример использования в модели

Создадим простую нейронную сеть с синхронизированной пакетной нормализацией:

import torch import torch.nn as nn class Model(nn.Module): def __init__(self): super().__init__() self.conv = nn.Conv2d(3, 16, 3) self.bn = nn.SyncBatchNorm(16) self.relu = nn.ReLU() def forward(self, x): x = self.conv(x) x = self.bn(x) return self.relu(x) model = Model() t = torch.randn(8, 3, 32, 32) res = model(t) print(res.shape)

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

torch.Size([8, 16, 30, 30])

Преобразование BatchNorm в SyncBatchNorm

Для упрощения перехода от обычной пакетной нормализации к синхронизированной в модуле предусмотрен метод convert_sync_batchnorm:

import torch import torch.nn as nn model = nn.Sequential( nn.Conv2d(3, 8, 3), nn.BatchNorm2d(8), nn.ReLU() ) sync_model = nn.SyncBatchNorm.convert_sync_batchnorm(model) print(type(sync_model[1]))

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

<class 'torch.nn.modules.batchnorm.SyncBatchNorm'>

Работа с атрибутами класса

Класс содержит основные атрибуты для управления состоянием нормализации:

import torch import torch.nn as nn bn = nn.SyncBatchNorm(5) print(bn.weight.shape) print(bn.bias.shape) print(bn.running_mean.shape) print(bn.running_var.shape) print("affine:", bn.affine) print("eps:", bn.eps)

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

torch.Size([5]) torch.Size([5]) torch.Size([5]) torch.Size([5]) affine: True eps: 1e-05

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

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