Spec-Zone.ru › PyTorch 2

Поддержка PyTorch 2.0 NNModule

Автор: Уильям Констанбл

torch.compile имеет специальную обработку для объектов torch.nn.Module, отслеживая их по-другому, чем произвольные классы Python, с целью создания более быстрого кода за счет предположений о структуре.

Этот документ описывает некоторые компромиссы или граничные случаи, возникающие из-за этой специализации.

Поддержка хуков NNModule

Ранее torch.compile не поддерживала хуки для nn.Modules, и если хуки регистрировались, они просто игнорировались в скомпилированной программе. Действительно, многие пользователи вообще не используют хуки nn.Module или используют их только для отладки, но существуют обоснованные случаи использования хуков nn.Module с torch.compile.

Хуки, которые организуются через реализацию nn.Module.__call__, включают _forward_pre_hooks, forward_hooks, _backward_pre_hooks, и _backward_hooks, и будут называться «хуками вызова». Эти хуки частично поддерживаются torch.compile с ограничениями, описанными ниже.

Другая категория хуков включает _state_dict_hooks и ее варианты pre и load_, и они все еще не поддерживаются torch.compile.

nn.Module.__call__ Использование и ограничения хуков

По умолчанию torch.compile будет отслеживать содержимое nn.Module.__call__, что означает, что он встретит и выполнит хуки forward/pre-forward. Если вы установите хуки перед вызовом torch.compile и затем не удалите или не измените хуки позже, ваш случай использования должен быть поддержан по умолчанию.

Хуки backward/Pre-backward, как правило, также поддерживаются, со схожими оговорками: в настоящее время разрывы графа в Dynamo происходят при доступе к словарям backward_hooks, что, вероятно, можно избежать при некоторых усилиях. Разрывы графа также влияют на время срабатывания хуков backward, так как сегменты графа выполняются как функции autograd, которые производят все свои градиенты одновременно. Предполагая, что Dynamo может не прерывать граф при наличии backward-хуков, мы по-прежнему ожидаем, что все backward хуки для ряда модулей сработают вместе после выполнения обратного прохода всего скомпилированного графа.

хуки на «разрешенных модулях» torch.compile обрабатывает такие распространенные модули, как torch.conv, а также модули, которые трудно отследить, особенно позволяя им вызываться непрозрачно в графе Dynamo вместо отслеживания Dynamo. Для таких модулей хуки в настоящее время вызывают разрыв графа, чтобы соответствующие модули работали вне Dynamo. В зависимости от модели это может привести к значительному снижению производительности, и требуются дополнительные усилия для улучшения этой поддержки.

skip_nnmodule_hook_guards По умолчанию torch._dynamo.config.skip_nnmodule_hook_guards установлено в True, что означает, что никакие стражи не будут установлены на каждом словаре хуков nn.Module, улучшая время выполнения за счет сокращения времени выполнения стражей, ценой неспособности заметить, если какой-либо словарь хуков изменяется после компиляции.

Если вы хотите иметь возможность удалять или изменять хуки после компиляции и иметь torch.compile реагировать соответствующим образом (перекомпилировать), вам необходимо установить skip_nnmodule_hook_guards=False и ожидать штрафа за время выполнения за добавленные стражи.

TODO: подтвердить, работают ли хуки backward/pre_backward или нет, и соответствующим образом задокументировать это

Хуки state_dict

Хуки state dict пока не поддерживаются в torch.compile.

TODO: выводить предупреждение warn_once при разрыве графа на хуках. Выводить warn_once, чтобы указать на этот документ, если хуки присутствуют.

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

Spec-Zone.ru

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