Spec-Zone.ru › PyTorch 2.14

Патчинг Batch Norm

Создано: 3 января 2023 г. | Последнее обновление: 11 июня 2025 г.

Что происходит?

Batch Norm требует обновления на месте running_mean и running_var, размер которых должен совпадать с размером входных данных. Functorch не поддерживает обновление на месте обычного тензора, принимающего пакетный тензор (т. е. regular.add_(batched) недопустимо). Поэтому при применении vmap к пакету входных данных для одного модуля возникает следующая ошибка

Как исправить

Один из наиболее надёжных способов — заменить BatchNorm на GroupNorm. Этот вариант поддерживают варианты 1 и 2

Все эти варианты предполагают, что статистика running вам не нужна. Если вы используете модуль, это означает, что предполагается, что вы не будете использовать batch norm в режиме оценки. Если у вас есть сценарий использования batch norm с vmap в режиме оценки, пожалуйста, сообщите об этом в задаче

Вариант 1: изменение BatchNorm

Чтобы заменить BatchNorm на GroupNorm, замените каждый экземпляр BatchNorm на:

BatchNorm2d(C, G, track_running_stats=False)

Здесь C — это тот же C, что и в исходном BatchNorm. G — число групп, на которые разбивается C. Таким образом, C % G == 0, а в качестве запасного варианта можно задать C == G, то есть каждый канал будет обрабатываться отдельно.

Если вам обязательно нужно использовать BatchNorm и вы самостоятельно создали модуль, можно изменить модуль так, чтобы он не использовал статистику running. Иными словами, для каждого модуля BatchNorm установите флаг track_running_stats в значение False

BatchNorm2d(64, track_running_stats=False)

Вариант 2: параметр torchvision

Некоторые модели torchvision, например resnet и regnet, могут принимать параметр norm_layer. Если для них задано значение по умолчанию, обычно это BatchNorm2d.

Вместо этого можно задать GroupNorm.

import torchvision
from functools import partial
torchvision.models.resnet18(norm_layer=lambda c: GroupNorm(num_groups=g, c))

Здесь, как и прежде, c % g == 0, поэтому в качестве запасного варианта задайте g = c.

Если вы предпочитаете BatchNorm, обязательно используйте версию, которая не использует статистику running

import torchvision
from functools import partial
torchvision.models.resnet18(norm_layer=partial(BatchNorm2d, track_running_stats=False))

Вариант 3: патчинг в functorch

В functorch добавлена функциональность, позволяющая быстро изменять модуль на месте, чтобы он не использовал статистику running. Замена слоя нормализации менее надёжна, поэтому мы не предлагаем такой возможности. Если у вас есть сеть, в которой BatchNorm не должен использовать статистику running, выполните replace_all_batch_norm_modules_, чтобы изменить модуль на месте и отключить использование статистики running

from torch.func import replace_all_batch_norm_modules_
replace_all_batch_norm_modules_(net)

Вариант 4: режим оценки

В режиме оценки running_mean и running_var не обновляются. Поэтому vmap поддерживает этот режим

model.eval()
vmap(model)(x)
model.train()

© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/func.batch_norm.html

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API