torch.jit.fork
-
torch.jit.fork(func, *args, **kwargs)[source] -
Создаёт асинхронную задачу, выполняющую
funcи ссылку на значение результата этого выполнения.forkвернётся немедленно, поэтому возвращаемое значениеfuncможет ещё не быть вычислено. Чтобы принудительно завершить задачу и получить возвращаемое значение, вызовитеtorch.jit.waitдля Future.fork, вызванный сfunc, возвращающийT, типизируется какtorch.jit.Future[T].forkвызовы могут быть произвольно вложенными и вызываться с позиционными и именованными аргументами. Асинхронное выполнение будет происходить только при запуске в TorchScript. Если запуск происходит в чистом Python,forkне будет выполняться параллельно.forkтакже не будет выполняться параллельно при вызове во время трассировки, однакоforkиwaitвызовы будут захвачены в экспортированном графе IR.Предупреждение
forkзадачи будут выполняться не детерминированно. Мы рекомендуем запускать только параллельные задачи fork для чистых функций, которые не изменяют их входные данные, атрибуты модулей или глобальное состояние.- Параметры
-
-
func (callable или torch.nn.Module) – Функция Python или
torch.nn.Module, которая будет вызвана. Если выполняется в TorchScript, она будет выполняться асинхронно, в противном случае – нет. Вызовы fork, подвергнутые трассировке, будут захвачены в IR. -
*args – аргументы для вызова
func. -
**kwargs – аргументы для вызова
func.
-
func (callable или torch.nn.Module) – Функция Python или
- Возвращаемое значение
-
ссылка на выполнение
func. ЗначениеTможет быть получено только принудительно завершивfuncчерезtorch.jit.wait. - Тип возвращаемого значения
-
torch.jit.Future[T]
Пример (fork свободной функции):
import torch from torch import Tensor def foo(a : Tensor, b : int) -> Tensor: return a + b def bar(a): fut : torch.jit.Future[Tensor] = torch.jit.fork(foo, a, b=2) return torch.jit.wait(fut) script_bar = torch.jit.script(bar) input = torch.tensor(2) # only the scripted version executes asynchronously assert script_bar(input) == bar(input) # trace is not run asynchronously, but fork is captured in IR graph = torch.jit.trace(bar, (input,)).graph assert "fork" in str(graph)Пример (fork метода модуля):
import torch from torch import Tensor class AddMod(torch.nn.Module): def forward(self, a: Tensor, b : int): return a + b class Mod(torch.nn.Module): def __init__(self): super(self).__init__() self.mod = AddMod() def forward(self, input): fut = torch.jit.fork(self.mod, a, b=2) return torch.jit.wait(fut) input = torch.tensor(2) mod = Mod() assert mod(input) == torch.jit.script(mod).forward(input)
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.jit.fork.html