Класс 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,
который применяет групповую нормализацию вместо пакетной