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