torch.jit.script
-
torch.jit.script(obj, optimize=None, _frames_up=0, _rcb=None, example_inputs=None)[source] -
Компиляция функции или
nn.Moduleбудет анализировать исходный код, компилировать его как код TorchScript с помощью компилятора TorchScript и возвращатьScriptModuleилиScriptFunction. Сам TorchScript — это подмножество языка Python, поэтому не все возможности Python работают, но мы предоставляем достаточно функционала для вычислений с тензорами и выполнения операций, зависящих от управления потоком. Для получения полного руководства обратитесь к Справочнику по языку TorchScript.Компиляция словаря или списка копирует данные внутри него в экземпляр TorchScript, который затем можно передавать по ссылке между Python и TorchScript без дополнительной копирования.
-
torch.jit.script can be used as a function for modules, functions, dictionaries and lists -
и как декоратор
@torch.jit.scriptдля Классов TorchScript и функций.
- Параметры:
-
-
obj (Callable, class, or nn.Module) – Компилируемый
nn.Module, функция, тип класса, словарь или список. -
example_inputs (Union[List[Tuple], Dict[Callable, List[Tuple]], None]) – Предоставьте примеры входных данных для аннотирования аргументов функции или
nn.Module.
-
obj (Callable, class, or nn.Module) – Компилируемый
- Возвращаемые значения:
-
Если
objявляетсяnn.Module, тоscriptвозвращает объектScriptModule. ВозвращаемыйScriptModuleбудет иметь тот же набор подмодулей и параметров, что и исходныйnn.Module. Еслиobj— это автономная функция, то будет возвращенScriptFunction. Еслиobjявляетсяdict, тоscriptвозвращает экземплярtorch._C.ScriptDict. Еслиobjявляетсяlist, тоscriptвозвращает экземплярtorch._C.ScriptList.
- Компиляция функции
-
Декоратор
@torch.jit.scriptсоздастScriptFunctionпутем компиляции тела функции.Пример (компиляция функции):
import torch @torch.jit.script def foo(x, y): if x.max() > y.max(): r = x else: r = y return r print(type(foo)) # torch.jit.ScriptFunction # See the compiled graph as Python code print(foo.code) # Call the function using the TorchScript interpreter foo(torch.ones(2, 2), torch.ones(2, 2)) - **Компиляция функции с использованием example_inputs
-
Примеры входных данных могут использоваться для аннотации аргументов функции.
Пример (аннотирование функции перед компиляцией):
import torch def test_sum(a, b): return a + b # Annotate the arguments to be int scripted_fn = torch.jit.script(test_sum, example_inputs=[(3, 4)]) print(type(scripted_fn)) # torch.jit.ScriptFunction # See the compiled graph as Python code print(scripted_fn.code) # Call the function using the TorchScript interpreter scripted_fn(20, 100) - Компиляция nn.Module
-
Компиляция
nn.Moduleпо умолчанию будет компилировать методforwardи рекурсивно компилировать любые методы, подмодули и функции, вызываемыеforward. Еслиnn.Moduleиспользует только функции, поддерживаемые TorchScript, то изменения в исходном коде модуля не потребуются.scriptсоздастScriptModuleс копиями атрибутов, параметров и методов исходного модуля.Пример (компиляция простого модуля с параметром):
import torch class MyModule(torch.nn.Module): def __init__(self, N, M): super(MyModule, self).__init__() # This parameter will be copied to the new ScriptModule self.weight = torch.nn.Parameter(torch.rand(N, M)) # When this submodule is used, it will be compiled self.linear = torch.nn.Linear(N, M) def forward(self, input): output = self.weight.mv(input) # This calls the `forward` method of the `nn.Linear` module, which will # cause the `self.linear` submodule to be compiled to a `ScriptModule` here output = self.linear(output) return output scripted_module = torch.jit.script(MyModule(2, 3))Пример (компиляция модуля с отслеженными подмодулями):
import torch import torch.nn as nn import torch.nn.functional as F class MyModule(nn.Module): def __init__(self): super(MyModule, self).__init__() # torch.jit.trace produces a ScriptModule's conv1 and conv2 self.conv1 = torch.jit.trace(nn.Conv2d(1, 20, 5), torch.rand(1, 1, 16, 16)) self.conv2 = torch.jit.trace(nn.Conv2d(20, 20, 5), torch.rand(1, 20, 16, 16)) def forward(self, input): input = F.relu(self.conv1(input)) input = F.relu(self.conv2(input)) return input scripted_module = torch.jit.script(MyModule())Для компиляции метода, отличного от
forward(и рекурсивной компиляции всего, что он вызывает), добавьте декоратор@torch.jit.exportк методу. Для отключения компиляции используйте@torch.jit.ignoreили@torch.jit.unused.Пример (экспортированный и проигнорированный метод в модуле):
import torch import torch.nn as nn class MyModule(nn.Module): def __init__(self): super(MyModule, self).__init__() @torch.jit.export def some_entry_point(self, input): return input + 10 @torch.jit.ignore def python_only_fn(self, input): # This function won't be compiled, so any # Python APIs can be used import pdb pdb.set_trace() def forward(self, input): if self.training: self.python_only_fn(input) return input * 99 scripted_module = torch.jit.script(MyModule()) print(scripted_module.some_entry_point(torch.randn(2, 2))) print(scripted_module(torch.randn(2, 2)))Пример (аннотирование forward nn.Module с помощью example_inputs):
import torch import torch.nn as nn from typing import NamedTuple class MyModule(NamedTuple): result: List[int] class TestNNModule(torch.nn.Module): def forward(self, a) -> MyModule: result = MyModule(result=a) return result pdt_model = TestNNModule() # Runs the pdt_model in eager model with the inputs provided and annotates the arguments of forward scripted_model = torch.jit.script(pdt_model, example_inputs={pdt_model: [([10, 20, ], ), ], }) # Run the scripted_model with actual inputs print(scripted_model([20]))
-
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.jit.script.html