Класс ScriptFunction
Класс ScriptFunction представляет собой скомпилированную версию функции Python, полученную с помощью декоратора @torch.jit.script. Объект этого класса содержит сериализованный граф вычислений, который может быть выполнен без доступа к исходному коду Python, что обеспечивает более высокую производительность и возможность сохранения на диск. Основное применение класса ScriptFunction - это создание моделей и функций в формате TorchScript для последующего использования в production-среде.
Синтаксис
import torch
@torch.jit.script
def function_name(params):
# тело функции с операциями PyTorch
return result
Пример
Давайте создадим простую скомпилированную функцию, которая вычисляет сумму квадратов элементов тензора:
import torch
@torch.jit.script
def sum_of_squares(x):
return torch.sum(x * x)
t = torch.tensor([1, 2, 3, 4, 5])
res = sum_of_squares(t)
print(res)
Результат выполнения кода:
tensor(55)
Пример
Объект ScriptFunction можно сохранить на диск и загрузить в другом окружении:
import torch
@torch.jit.script
def add_tensors(a, b):
return a + b
t1 = torch.tensor([1, 2, 3])
t2 = torch.tensor([4, 5, 6])
res = add_tensors(t1, t2)
print(res)
torch.jit.save(add_tensors, 'add_tensors.pt')
Результат выполнения кода:
tensor([5, 7, 9])
Пример
Загруженный объект можно использовать как обычную функцию:
import torch
loaded_func = torch.jit.load('add_tensors.pt')
t3 = torch.tensor([10, 20, 30])
t4 = torch.tensor([1, 2, 3])
res = loaded_func(t3, t4)
print(res)
Результат выполнения кода:
tensor([11, 22, 33])
Пример
Скомпилированная функция поддерживает использование условных операторов и циклов:
import torch
@torch.jit.script
def conditional_add(a, b, condition):
if condition:
return a + b
else:
return a - b
t1 = torch.tensor([10, 20, 30])
t2 = torch.tensor([1, 2, 3])
cond = True
res = conditional_add(t1, t2, cond)
print(res)
Результат выполнения кода:
tensor([11, 22, 33])
Смотрите также
-
функцию
jit.script,
которая преобразует функцию Python в формат TorchScript -
функцию
jit.save,
которая сохраняет объект TorchScript на диск -
функцию
jit.load,
которая загружает объект TorchScript с диска -
класс
RecursiveScriptModule,
который представляет скомпилированный модуль TorchScript