РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
632 of 769 menu

Функция 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-модуль для инференса
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить