Spec-Zone.ru › PyTorch 2

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

Spec-Zone.ru

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