Патчинг 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