torch.xpu.make_graphed_callables
-
torch.xpu.make_graphed_callables(callables: Module | Callable[[...], object], sample_args: tuple[Tensor, ...], num_warmup_iters: int = 3, allow_unused_input: bool = False, pool: _POOL_HANDLE | None = None) → Module | Callable[[...], object][source] - torch.xpu.make_graphed_callables(callables:tuple[Module|Callable[[...],object],...], sample_args:tuple[tuple[Tensor,...],...], num_warmup_iters:int=3, allow_unused_input:bool=False, pool:_POOL_HANDLE|None=None) tuple[Module|Callable[[...],object],...]
-
Принимает вызываемые объекты (функции или
nn.Module) и возвращает их версии, захваченные в граф.Прямой проход каждого вызываемого объекта, захваченного в граф, выполняет работу XPU прямого прохода исходного вызываемого объекта в виде графа XPU внутри одного узла autograd.
Прямой проход вызываемого объекта, захваченного в граф, также добавляет узел обратного прохода в граф autograd. Во время обратного прохода этот узел выполняет работу XPU обратного прохода вызываемого объекта в виде графа XPU.
Таким образом, каждый вызываемый объект, захваченный в граф, должен быть прямой заменой исходного вызываемого объекта в цикле обучения с включённым autograd.
Подробные сведения об использовании и ограничениях см. в разделе Частичный захват сети.
Если передать кортеж из нескольких вызываемых объектов, для их захвата будет использоваться один и тот же пул памяти.
- Параметры:
-
- callables (torch.nn.Module или функция Python, или tuple из таких объектов) – Вызываемый объект или объекты для захвата в граф. Если передан кортеж вызываемых объектов, их порядок в кортеже должен совпадать с порядком их выполнения в рабочей нагрузке.
-
sample_args (tuple тензоров или tuple из кортежей тензоров) – Примеры аргументов для каждого вызываемого объекта. Если передан один вызываемый объект,
sample_argsдолжен быть одним кортежем тензоров-аргументов. Если передан кортеж вызываемых объектов,sample_argsдолжен быть кортежем кортежей тензоров-аргументов. -
num_warmup_iters (int) – Количество итераций прогрева. В настоящее время для прогрева
DataDistributedParallelтребуется 11 итераций. Значение по умолчанию:3. - allow_unused_input (bool) – Если False, указание входных данных, которые не использовались при вычислении выходных данных (и градиент которых поэтому всегда равен нулю), считается ошибкой. Значение по умолчанию — False.
-
pool (необязательный) – Токен (возвращаемый функцией
graph_pool_handle()илиother_Graph_instance.pool()), указывающий, что этот граф может совместно использовать память с указанным пулом.
Примечание
Состояние
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.amp.autocast()должен иметьcache_enabled=False.
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.xpu.make_graphed_callables.html