Spec-Zone.ru › PyTorch 1

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. Значение 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

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API