torch.cuda.make_graphed_callables
-
torch.cuda.make_graphed_callables(callables, sample_args, num_warmup_iters=3, allow_unused_input=False)[source] -
Принимает вызываемые объекты (функции или
nn.Module) и возвращает их графовые версии.Каждый графовый вызываемый объект в ходе прямого прохода выполняет исходную вычислительную работу вызываемого объекта в CUDA как CUDA-граф внутри одного узла autograd.
Прямой проход графового вызываемого объекта также добавляет обратный узел в граф autograd. Во время обратного прохода этот узел выполняет обратную работу вызываемого объекта как CUDA-граф.
Таким образом, каждый графовый вызываемый объект должен быть прямым заменителем исходного вызываемого объекта в цикле обучения, поддерживающем autograd.
Подробное использование и ограничения см. в Частичном захвате сети.
Если вы передаете кортеж из нескольких вызываемых объектов, их захваты будут использовать один и тот же пул памяти. См. Управление памятью графа, чтобы понять, когда это уместно.
- Parameters
-
- callables (torch.nn.Module или Python-функция, или кортеж из этих) – Вызываемый объект или вызываемые объекты для построения графа. См. Управление памятью графа, чтобы понять, когда уместно передавать кортеж вызываемых объектов. Если вы передаете кортеж вызываемых объектов, их порядок в кортеже должен соответствовать порядку их выполнения в реальной рабочей нагрузке.
-
sample_args (кортеж из тензоров, или кортеж из кортежей из тензоров) – Образцы аргументов для каждого вызываемого объекта. Если был передан один вызываемый объект,
sample_argsдолжен быть единственным кортежем тензоров аргументов. Если был передан кортеж вызываемых объектов,sample_argsдолжен быть кортежем кортежей тензоров аргументов. -
num_warmup_iters (int) – Количество итераций разогрева. В настоящее время
DataDistributedParallelтребуется 11 итераций для разогрева. По умолчанию:3. - allow_unused_input (bool) – Если False, указание входов, которые не использовались при вычислении выходов (и, следовательно, их градиент всегда равен нулю), является ошибкой. По умолчанию False.
Примечание
Состояние каждого тензора в
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/2.1/generated/torch.cuda.make_graphed_callables.html