Замена 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