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/1.13/generated/torch.jit.fork.html