Функция jit.script
Функция jit.script преобразует Python-функцию или класс модуля в TorchScript.
Это статически типизированное подмножество Python, которое может быть оптимизировано и выполнено вне среды Python.
Первым параметром функция принимает вызываемый объект (функцию или класс модуля).
Вторым параметром можно передать дополнительные опции компиляции через torch.jit.ScriptFunction или torch.jit.RecursiveScriptModule.
Синтаксис
torch.jit.script(func_or_module)
Пример
Давайте преобразуем простую функцию, которая вычисляет сумму двух тензоров:
import torch
@torch.jit.script
def sum_tensors(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
return a + b
t1 = torch.tensor([1, 2, 3])
t2 = torch.tensor([4, 5, 6])
res = sum_tensors(t1, t2)
print(res)
Результат выполнения кода:
tensor([5, 7, 9])
Пример
Преобразуем класс модуля nn.Module с помощью декоратора:
import torch
import torch.nn as nn
@torch.jit.script
class MyModule(nn.Module):
def __init__(self):
super().__init__()
self.linear = nn.Linear(2, 2)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.linear(x)
model = MyModule()
x = torch.tensor([[1.0, 2.0]])
res = model(x)
print(res)
Результат выполнения кода:
tensor([[-0.2940, -0.2592]], grad_fn=<AddmmBackward0>)
Пример
Используем jit.script как функцию для преобразования существующего модуля:
import torch
import torch.nn as nn
class SimpleNet(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(3, 1)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.fc(x)
model = SimpleNet()
scripted_model = torch.jit.script(model)
x = torch.tensor([[1.0, 2.0, 3.0]])
res = scripted_model(x)
print(res)
Результат выполнения кода:
tensor([[-0.2198]], grad_fn=<AddmmBackward0>)
Смотрите также
-
функцию
jit.trace,
которая создает TorchScript путем трассировки примера выполнения -
функцию
jit.save,
которая сохраняет TorchScript модуль на диск -
функцию
jit.load,
которая загружает TorchScript модуль с диска -
функцию
jit.freeze,
которая замораживает модуль, оптимизируя его для инференса