Функция jit.script_if_tracing
Функция jit.script_if_tracing применяется для декорирования функций, которые должны компилироваться в TorchScript только в том случае, если код выполняется внутри контекста трассировки (например, при использовании torch.jit.trace). В обычном режиме выполнения такая функция остаётся обычной Python-функцией. Это полезно для оптимизации производительности и совместимости с трассировкой.
Параметры функции:
-
fn- декорируемая функция или метод, который будет скомпилирован в TorchScript при трассировке. -
_checks- внутренний параметр, управляющий проверками (в основном используется в самом PyTorch).
Функция возвращает декоратор, который при трассировке заменяет исходную функцию на её скомпилированную версию.
Синтаксис
torch.jit.script_if_tracing(fn)
Пример
Давайте рассмотрим базовое использование декоратора script_if_tracing для условной компиляции функции:
import torch
@torch.jit.script_if_tracing
def conditional_fn(x):
return x + 1
t = torch.tensor([1, 2, 3, 4, 5])
res = conditional_fn(t)
print(res)
Результат выполнения кода:
tensor([2, 3, 4, 5, 6])
Пример
В этом примере функция компилируется только при трассировке модели, а в обычном режиме остаётся Python-функцией:
import torch
@torch.jit.script_if_tracing
def traced_only_fn(x, y):
return x * y
class MyModule(torch.nn.Module):
def forward(self, x):
return traced_only_fn(x, x)
model = MyModule()
t = torch.tensor([1, 2, 3, 4, 5])
traced_model = torch.jit.trace(model, t)
print(traced_model(t))
Результат выполнения кода:
tensor([1, 4, 9, 16, 25])
Пример
Покажем, как script_if_tracing влияет на выполнение в обычном режиме и внутри трассировки. В обычном режиме функция выполняется как обычный Python-код:
import torch
@torch.jit.script_if_tracing
def smart_fn(x):
return x + 10
t = torch.tensor([1, 2, 3, 4, 5])
res = smart_fn(t)
print(res)
Результат выполнения кода:
tensor([11, 12, 13, 14, 15])
Пример
Используем script_if_tracing внутри метода модели, чтобы код компилировался только при трассировке:
import torch
class Net(torch.nn.Module):
def __init__(self):
super().__init__()
self.linear = torch.nn.Linear(5, 3)
@torch.jit.script_if_tracing
def inner_func(self, x):
return x + 1
def forward(self, x):
x = self.inner_func(x)
return self.linear(x)
model = Net()
t = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0])
traced = torch.jit.trace(model, t)
res = traced(t)
print(res)
Результат выполнения кода:
tensor([1.9319, 2.0698, 2.9400], grad_fn=<AddBackward0>)
Смотрите также
-
функцию
jit.script,
которая компилирует функцию или класс в TorchScript -
функцию
jit.trace,
которая создаёт TorchScript-модуль путём трассировки -
функцию
jit.freeze,
которая замораживает TorchScript-модуль для инференса -
функцию
jit.optimize_for_inference,
которая оптимизирует TorchScript-модуль для инференса