Spec-Zone.ru › PyTorch 1

torch.cuda.make_graphed_callables

torch.cuda.make_graphed_callables(callables, sample_args, num_warmup_iters=3) [source]

Принимает вызываемые объекты (функции или nn.Module) и возвращает их графические версии.

Каждый графический вызываемый объект в своей прямой передаче выполняет исходную работу вызываемого объекта CUDA как граф CUDA внутри одного узла autograd.

Прямая передача графического вызываемого объекта также добавляет обратный узел в граф autograd. Во время обратного прохода этот узел выполняет обратную работу вызываемого объекта как граф CUDA.

Поэтому каждый графический вызываемый объект должен быть полным аналогом своего исходного вызываемого объекта в цикле обучения с поддержкой autograd.

Подробное использование и ограничения см. в Частичном захвате сети.

Если вы передаете кортеж из нескольких вызываемых объектов, их захват будет использовать один и тот же пул памяти. Когда это уместно, см. Управление памятью графа.

Параметры:
  • callables (torch.nn.Module или Python-функция или кортеж из них) – Вызываемый объект или вызываемые объекты для построения графа. Когда уместно передавать кортеж вызываемых объектов, см. Управление памятью графа. Если вы передаете кортеж вызываемых объектов, их порядок в кортеже должен соответствовать порядку их выполнения в реальной рабочей нагрузке.
  • sample_args (кортеж тензоров или кортеж кортежей тензоров) – Образцы аргументов для каждого вызываемого объекта. Если был передан один вызываемый объект, sample_args должен быть единственным кортежем тензоров аргументов. Если был передан кортеж вызываемых объектов, sample_args должен быть кортежем кортежей тензоров аргументов.
  • num_warmup_iters (int) – Количество итераций разогрева. В настоящее время DataDistributedParallel требуется 11 итераций для разогрева. По умолчанию: 3.

Примечание

Состояние requires_grad каждого тензора в sample_args должно соответствовать ожидаемому состоянию соответствующего реального входного значения в цикле обучения.

Предупреждение

Этот API находится в стадии бета-тестирования и может быть изменён в будущих выпусках.

Предупреждение

sample_args для каждого вызываемого объекта должен быть кортежем тензоров. Другие типы и ключевые аргументы недопустимы.

Предупреждение

Возвращаемые вызываемые объекты не поддерживают дифференцирование более высокого порядка (например, двойное обратное дифференцирование).

Предупреждение

В любом Module переданном в make_graphed_callables(), только параметры могут быть обучены. Буферы должны иметь requires_grad=False.

Предупреждение

После передачи torch.nn.Module через make_graphed_callables(), вы не можете добавлять или удалять параметры или буферы данного модуля.

Предупреждение

torch.nn.Module, переданные в make_graphed_callables(), не должны иметь зарегистрированных хуков модуля на момент их передачи. Однако регистрация хуков на модулях после их передачи через make_graphed_callables() разрешена.

Предупреждение

При выполнении графического вызываемого объекта аргументы должны передаваться в том же порядке и формате, в котором они появились в sample_args этого вызываемого объекта.

Предупреждение

Автоматическое смешанное точность поддерживается в make_graphed_callables() только при отключенном кэшировании. У контекстного менеджера torch.cuda.amp.autocast() должен быть cache_enabled=False.

Предупреждение

Все выходные тензоры графических вызываемых объектов должны требовать градиента.

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

Spec-Zone.ru

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