Spec-Zone.ru › PyTorch 2

Замена Batch Norm

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

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

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

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

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

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

Если вы хотите изменить на GroupNorm, везде, где у вас есть BatchNorm, замените его на:

BatchNorm2d(C, G, track_running_stats=False)

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

Если вам необходимо использовать BatchNorm, и вы сами создали модуль, вы можете изменить модуль так, чтобы он не использовал running stats. Другими словами, везде, где есть модуль 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 stats

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

Вариант 3: Обновление с помощью functorch

Functorch добавил некоторую функциональность, позволяющую быстро и на месте обновить модуль так, чтобы он не использовал running stats. Изменение слоя нормализации более хрупко, поэтому мы не предлагаем это. Если у вас есть сеть, где вы хотите, чтобы BatchNorm не использовал running stats, вы можете запустить replace_all_batch_norm_modules_ для обновления модуля на месте так, чтобы он не использовал running stats

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()

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

Spec-Zone.ru

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