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

Класс GroupNorm

Класс GroupNorm применяет групповую нормализацию к входным данным. Основная идея метода заключается в разделении каналов на группы и вычислении среднего и дисперсии внутри каждой группы независимо. Это позволяет нормализовать активации в пределах группы каналов, что особенно полезно для моделей с ограниченным размером батча. Первым параметром передаётся количество групп num_groups, вторым - количество каналов num_channels.

Синтаксис

torch.nn.GroupNorm(num_groups, num_channels, eps=1e-5, affine=True)

Параметры:

num_groups - количество групп, на которые делятся каналы (должно делиться на число каналов);

num_channels - количество каналов во входных данных;

eps - значение для численной стабильности (по умолчанию 1e-5);

affine - булевый параметр, указывающий, следует ли использовать обучаемые параметры сдвига и масштаба (по умолчанию True).

Пример

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

import torch import torch.nn as nn # Создаем слой GroupNorm: 2 группы, 4 канала gn = nn.GroupNorm(num_groups=2, num_channels=4) # Входной тензор: батч размером 1, 4 канала, 3x3 t = torch.randn(1, 4, 3, 3) res = gn(t) print(res.shape)

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

torch.Size([1, 4, 3, 3])

Пример

Применим групповую нормализацию к трёхмерному тензору с явным указанием параметров:

import torch import torch.nn as nn torch.manual_seed(0) # Создаем слой GroupNorm с 3 группами и 6 каналами gn = nn.GroupNorm(num_groups=3, num_channels=6, eps=1e-6, affine=True) # Входной тензор размером 2x6x4x4 t = torch.randn(2, 6, 4, 4) res = gn(t) print(res[0, 0, 0, :5])

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

tensor([-0.2377, 0.8390, -1.5702, 1.1739, -0.1737])

Пример

Рассмотрим, как GroupNorm работает с одномерными данными:

import torch import torch.nn as nn # Слой GroupNorm для одномерных данных gn = nn.GroupNorm(num_groups=2, num_channels=4) # Входной тензор: батч 3, 4 канала, 10 элементов t = torch.randn(3, 4, 10) res = gn(t) print(res.shape) print(t.mean(dim=(2,), keepdim=True).shape)

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

torch.Size([3, 4, 10]) torch.Size([3, 4, 1])

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

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