Справочник по языку TorchScript
TorchScript — это статически типизированное подмножество Python, которое можно написать непосредственно (используя декоратор @torch.jit.script) или сгенерировать автоматически из кода Python с помощью трассировки. При использовании трассировки код автоматически преобразуется в это подмножество Python путём записи только фактических операций над тензорами и простого выполнения и отбрасывания другого окружающего кода Python.
При непосредственном написании TorchScript с помощью @torch.jit.script декоратора программист должен использовать только подмножество Python, поддерживаемое TorchScript. В этом разделе описывается, что поддерживается в TorchScript, как если бы это был справочник по отдельному языку программирования. Любые функции Python, не упомянутые в этом справочнике, не являются частью TorchScript. Обратитесь к Builtin Functions для получения полного справочника по доступным методам, модулям и функциям тензоров PyTorch.
Как подмножество Python, любая допустимая функция TorchScript также является допустимой функцией Python. Это позволяет disable TorchScript и отлаживать функцию с помощью стандартных инструментов Python, таких как pdb. Обратное неверно: существует много допустимых программ Python, которые не являются допустимыми программами TorchScript. Вместо этого TorchScript фокусируется конкретно на функциях Python, необходимых для представления нейронных сетей в PyTorch.
Типы
Главное отличие TorchScript от полного языка Python заключается в том, что TorchScript поддерживает только небольшой набор типов, необходимых для выражения моделей нейронных сетей. В частности, TorchScript поддерживает:
Тип | Описание |
|---|---|
| Тензор PyTorch любого типа, размерности или бэкенда |
| Кортеж, содержащий подтипы |
| Логическое значение |
| Целое число в виде скаляра |
| Скалярное число с плавающей точкой |
| Строка |
| Список, все члены которого имеют тип |
| Значение, которое либо равно None, либо имеет тип |
| Словарь с ключом типа |
| Класс TorchScript |
| Перечисление TorchScript |
|
|
| Один из подтипов |
В отличие от Python, каждая переменная в функции TorchScript должна иметь один статический тип. Это упрощает оптимизацию функций TorchScript.
Пример (несоответствие типов)
import torch
@torch.jit.script
def an_error(x):
if x:
r = torch.rand(1)
else:
r = 4
return r
Traceback (most recent call last):
...
RuntimeError: ...
Type mismatch: r is set to type Tensor in the true branch and type int in the false branch:
@torch.jit.script
def an_error(x):
if x:
~~~~~
r = torch.rand(1)
~~~~~~~~~~~~~~~~~
else:
~~~~~
r = 4
~~~~~ <--- HERE
return r
and was used here:
else:
r = 4
return r
~ <--- HERE...
Неподдерживаемые конструкции типизации
TorchScript не поддерживает все возможности и типы модуля typing. Некоторые из них являются более фундаментальными вещами, которые вряд ли будут добавлены в будущем, в то время как другие могут быть добавлены, если будет достаточно пользовательского спроса, чтобы сделать это приоритетным.
Эти типы и возможности модуля typing недоступны в TorchScript.
Элемент | Описание |
|---|---|
| |
Не реализовано | |
Не реализовано | |
Не реализовано | |
Не реализовано | |
Не реализовано | |
Поддерживается для аннотаций атрибутов модуля, но не для функций | |
TorchScript не поддерживает | |
| |
Псевдонимы типов | Не реализовано |
Номинальное и структурное субобразование | Номинальная типизация находится в разработке, но структурная типизация нет |
NewType | Вероятно, не будет реализовано |
Обобщения | Вероятно, не будет реализовано |
Любая другая функциональность из модуля typing, не указанная в этой документации, не поддерживается.
Типы по умолчанию
По умолчанию все параметры функции TorchScript предполагаются тензорами. Чтобы указать, что аргумент функции TorchScript имеет другой тип, можно использовать аннотации типа в стиле MyPy с использованием указанных выше типов.
import torch
@torch.jit.script
def foo(x, tup):
# type: (int, Tuple[Tensor, Tensor]) -> Tensor
t0, t1 = tup
return t0 + t1 + x
print(foo(3, (torch.rand(3), torch.rand(3))))
Примечание
Также можно аннотировать типы с помощью подсказок типов Python 3 из модуля typing.
import torch
from typing import Tuple
@torch.jit.script
def foo(x: int, tup: Tuple[torch.Tensor, torch.Tensor]) -> torch.Tensor:
t0, t1 = tup
return t0 + t1 + x
print(foo(3, (torch.rand(3), torch.rand(3))))
Пустой список предполагается типа List[Tensor], а пустые словари — типа Dict[str, Tensor]. Для создания пустых списков или словарей других типов используйте Python 3 type hints.
Пример (аннотации типов для Python 3):
import torch
import torch.nn as nn
from typing import Dict, List, Tuple
class EmptyDataStructures(torch.nn.Module):
def __init__(self):
super(EmptyDataStructures, self).__init__()
def forward(self, x: torch.Tensor) -> Tuple[List[Tuple[int, float]], Dict[str, int]]:
# This annotates the list to be a `List[Tuple[int, float]]`
my_list: List[Tuple[int, float]] = []
for i in range(10):
my_list.append((i, x.item()))
my_dict: Dict[str, int] = {}
return my_list, my_dict
x = torch.jit.script(EmptyDataStructures())
Необязательное уточнение типа
TorchScript уточнит тип переменной типа Optional[T] при сравнении с None внутри условия оператора if или при проверке в assert. Компилятор может обосновать несколько проверок None, которые объединены с and, or, и not . Уточнение также произойдет для блоков else операторов if, которые не явно написаны.
Проверка None должна находиться внутри условия оператора if; присвоение проверки None переменной и использование её в условии оператора if не будут уточнять типы переменных в проверке. Уточняются только локальные переменные, атрибут типа self.x не будет, и его необходимо присвоить локальной переменной для уточнения.
Пример (уточнение типов параметров и локальных переменных):
import torch
import torch.nn as nn
from typing import Optional
class M(nn.Module):
z: Optional[int]
def __init__(self, z):
super(M, self).__init__()
# If `z` is None, its type cannot be inferred, so it must
# be specified (above)
self.z = z
def forward(self, x, y, z):
# type: (Optional[int], Optional[int], Optional[int]) -> int
if x is None:
x = 1
x = x + 1
# Refinement for an attribute by assigning it to a local
z = self.z
if y is not None and z is not None:
x = y + z
# Refinement via an `assert`
assert z is not None
x += z
return x
module = torch.jit.script(M(2))
module = torch.jit.script(M(None))
Классы TorchScript
Предупреждение
Поддержка классов TorchScript находится в стадии эксперимента. В настоящее время она лучше всего подходит для простых типов записей (например, NamedTuple с прикрепленными методами).
Классы Python могут использоваться в TorchScript, если они аннотированы с помощью @torch.jit.script, аналогично тому, как вы объявляете функцию TorchScript:
@torch.jit.script
class Foo:
def __init__(self, x, y):
self.x = x
def aug_add_x(self, inc):
self.x += inc
Этот подмножество ограничено:
- Все функции должны быть допустимыми функциями TorchScript (включая
__init__()). - Классы должны быть классами нового стиля, так как мы используем
__new__()для их создания с помощью pybind11. -
Классы TorchScript являются статически типизированными. Члены могут быть объявлены только путем присвоения значения self в методе
__init__().Например, присвоение значений
selfвне метода__init__():@torch.jit.script class Foo: def assign_x(self): self.x = torch.rand(2, 3)Приведёт к:
RuntimeError: Tried to set nonexistent attribute: x. Did you forget to initialize it in __init__()?: def assign_x(self): self.x = torch.rand(2, 3) ~~~~~~~~~~~~~~~~~~~~~~~~ <--- HERE
- В теле класса не допускаются выражения, кроме определений методов.
- Не поддерживается наследование или любая другая стратегия полиморфизма, кроме наследования от
objectдля указания класса нового стиля.
После определения класса его можно использовать как в TorchScript, так и в Python взаимозаменяемо, как любой другой тип TorchScript:
# Declare a TorchScript class
@torch.jit.script
class Pair:
def __init__(self, first, second):
self.first = first
self.second = second
@torch.jit.script
def sum_pair(p):
# type: (Pair) -> Tensor
return p.first + p.second
p = Pair(torch.rand(2, 3), torch.rand(2, 3))
print(sum_pair(p))
Перечисления TorchScript
Перечисления Python могут использоваться в TorchScript без каких-либо дополнительных аннотаций или кода:
from enum import Enum
class Color(Enum):
RED = 1
GREEN = 2
@torch.jit.script
def enum_fn(x: Color, y: Color) -> bool:
if x == Color.RED:
return True
return x == y
После определения перечисления его можно использовать как в TorchScript, так и в Python взаимозаменяемо, как любой другой тип TorchScript. Тип значений перечисления должен быть int, float, или str . Все значения должны быть одного типа; неоднородные типы значений перечисления не поддерживаются.
Именованные кортежи
Типы, созданные с помощью collections.namedtuple, могут использоваться в TorchScript.
import torch
import collections
Point = collections.namedtuple('Point', ['x', 'y'])
@torch.jit.script
def total(point):
# type: (Point) -> Tensor
return point.x + point.y
p = Point(x=torch.rand(3), y=torch.rand(3))
print(total(p))
Итерируемые объекты
Некоторые функции (например, zip и enumerate) могут работать только с итерируемыми типами. Итерируемые типы в TorchScript включают списки, кортежи, словари, строки, Tensor и torch.nn.ModuleList.
Выражения
Поддерживаются следующие выражения Python.
Литералы
True False None 'string literals' "string literals" 3 # interpreted as int 3.4 # interpreted as a float
Создание списков
Пустой список предполагается иметь тип List[Tensor]. Типы других списков вычисляются на основе типа элементов. Подробнее см. Типы по умолчанию.
[3, 4] [] [torch.rand(3), torch.rand(4)]
Создание кортежей
(3, 4) (3,)
Создание словарей
Пустой словарь предполагается иметь тип Dict[str, Tensor]. Типы других словарей вычисляются на основе типа элементов. Подробнее см. Типы по умолчанию.
{'hello': 3}
{}
{'a': torch.rand(3), 'b': torch.rand(4)}
Переменные
Как разрешаются переменные, см. Разрешение переменных.
my_variable_name
Арифметические операторы
a + b a - b a * b a / b a ^ b a @ b
Операторы сравнения
a == b a != b a < b a > b a <= b a >= b
Логические операторы
a and b a or b not b
Индексы и срезы
t[0] t[-1] t[0:2] t[1:] t[:1] t[:] t[0, 1] t[0, 1:2] t[0, :1] t[-1, 1:, 0] t[1:, -1, 0] t[i:j, i]
Вызовы функций
Вызовы функции builtin functions
torch.rand(3, dtype=torch.int)
Вызовы других функций скрипта:
import torch
@torch.jit.script
def foo(x):
return x + 1
@torch.jit.script
def bar(x):
return foo(x)
Вызовы методов
Вызовы методов встроенных типов, таких как тензор: x.mm(y)
Для модулей методы необходимо скомпилировать перед вызовом. Компилятор TorchScript рекурсивно компилирует методы, которые он видит при компиляции других методов. По умолчанию компиляция начинается с метода forward. Любые методы, вызываемые forward, будут скомпилированы, а любые методы, вызываемые этими методами, и так далее. Чтобы начать компиляцию с другого метода, отличного от forward, используйте декоратор @torch.jit.export (forward неявным образом помечен как @torch.jit.export).
Прямой вызов подмодуля (например, self.resnet(input)) эквивалентен вызову его метода forward (например, self.resnet.forward(input)).
import torch
import torch.nn as nn
import torchvision
class MyModule(nn.Module):
def __init__(self):
super(MyModule, self).__init__()
means = torch.tensor([103.939, 116.779, 123.68])
self.means = torch.nn.Parameter(means.resize_(1, 3, 1, 1))
resnet = torchvision.models.resnet18()
self.resnet = torch.jit.trace(resnet, torch.rand(1, 3, 224, 224))
def helper(self, input):
return self.resnet(input - self.means)
def forward(self, input):
return self.helper(input)
# Since nothing in the model calls `top_level_method`, the compiler
# must be explicitly told to compile this method
@torch.jit.export
def top_level_method(self, input):
return self.other_helper(input)
def other_helper(self, input):
return input + 10
# `my_script_module` will have the compiled methods `forward`, `helper`,
# `top_level_method`, and `other_helper`
my_script_module = torch.jit.script(MyModule())
Тernary выражения
x if x > y else y
Преобразования типов
float(ten) int(3.5) bool(ten) str(2)``
Доступ к параметрам модуля
self.my_parameter self.my_submodule.my_parameter
Утверждения
TorchScript поддерживает следующие типы утверждений:
Простые присваивания
a = b a += b # short-hand for a = a + b, does not operate in-place on a a -= b
Присваивания с сопоставлением шаблонов
a, b = tuple_or_list a, b, *c = a_tuple
Множественные присваивания
a = b, c = tup
Утверждения печати
print("the result of an add:", a + b)
Утверждения if
if a < 4:
r = -a
elif a < 3:
r = a + a
else:
r = 3 * a
В дополнение к boolean, в условных операторах могут использоваться значения float, int и тензоры; они будут неявным образом преобразованы к типу boolean.
Циклы while
a = 0
while a < 4:
print(a)
a += 1
Циклы for с range
x = 0
for i in range(10):
x *= i
Циклы for по кортежам
Эти циклы раскручиваются, генерируя тело для каждого члена кортежа. Тело должно корректно проверять типы для каждого члена.
tup = (3, torch.rand(4))
for x in tup:
print(x)
Циклы for по константным nn.ModuleList
Чтобы использовать nn.ModuleList внутри скомпилированного метода, его нужно пометить как константный, добавив имя атрибута в список __constants__ для типа. Циклы for по nn.ModuleList будут раскручиваться во время компиляции, и каждое члену константного списка модулей будет соответствовать тело цикла.
class SubModule(torch.nn.Module):
def __init__(self):
super(SubModule, self).__init__()
self.weight = nn.Parameter(torch.randn(2))
def forward(self, input):
return self.weight + input
class MyModule(torch.nn.Module):
__constants__ = ['mods']
def __init__(self):
super(MyModule, self).__init__()
self.mods = torch.nn.ModuleList([SubModule() for i in range(10)])
def forward(self, v):
for module in self.mods:
v = module(v)
return v
m = torch.jit.script(MyModule())
Break и Continue
for i in range(5):
if i == 1:
continue
if i == 3:
break
print(i)
Возврат
return a, b
Разрешение переменных
TorchScript поддерживает подмножество правил разрешения переменных (т. е. области видимости) Python. Локальные переменные ведут себя так же, как в Python, за исключением ограничения, что переменная должна иметь один и тот же тип на всех путях функции. Если переменная имеет разные типы на разных ветвях оператора if, её использование после окончания оператора if является ошибкой.
Аналогично, переменная не может быть использована, если она определена только на некоторых путях функции.
Пример:
@torch.jit.script
def foo(x):
if x < 0:
y = 4
print(y)
Traceback (most recent call last):
...
RuntimeError: ...
y is not defined in the false branch...
@torch.jit.script...
def foo(x):
if x < 0:
~~~~~~~~~
y = 4
~~~~~ <--- HERE
print(y)
and was used here:
if x < 0:
y = 4
print(y)
~ <--- HERE...
Нелокальные переменные разрешаются до значений Python во время компиляции, когда функция определяется. Эти значения затем преобразуются в значения TorchScript с помощью правил, описанных в Использовании значений Python.
Использование значений Python
Для удобства написания TorchScript мы разрешаем коду скрипта ссылаться на значения Python в области видимости окружения. Например, каждый раз, когда есть ссылка на torch, компилятор TorchScript фактически разрешает её до модуля Python torch при объявлении функции. Эти значения Python не являются частью TorchScript в качестве первого класса. Вместо этого они десугариваются во время компиляции в примитивные типы, которые поддерживает TorchScript. Это зависит от динамического типа значения Python, на которое ссылаются во время компиляции. Этот раздел описывает правила, используемые при обращении к значениям Python в TorchScript.
Функции
TorchScript может вызывать функции Python. Эта функциональность очень полезна при постепенном преобразовании модели в TorchScript. Модель может быть перенесена в TorchScript функцию за функцией, оставляя вызовы функций Python на месте. Таким образом, вы можете поэтапно проверять корректность модели по мере продвижения.
-
torch.jit.is_scripting()[source] -
Функция, возвращающая True, когда выполняется компиляция, и False в противном случае. Это особенно полезно с декоратором @unused для оставления в вашей модели кода, который пока не совместим с TorchScript. .. testcode:
import torch @torch.jit.unused def unsupported_linear_op(x): return x def linear(x): if torch.jit.is_scripting(): return torch.linear(x) else: return unsupported_linear_op(x)- Тип возвращаемого значения:
-
torch.jit.is_tracing()[source] -
Возвращает
Trueво время трассировки (если функция вызывается во время трассировки кода сtorch.jit.trace) иFalseв противном случае.
Поиск атрибутов в модулях Python
TorchScript может искать атрибуты в модулях. Так осуществляется доступ к Builtin functions, таким как torch.add. Это позволяет TorchScript вызывать функции, определённые в других модулях.
Определённые в Python константы
TorchScript также предоставляет способ использования констант, определённых в Python. Их можно использовать для вставки гиперпараметров в функцию или определения универсальных констант. Есть два способа указать, что значение Python должно рассматриваться как константа.
- Значения, полученные в качестве атрибутов модуля, предполагаются константами:
import math
import torch
@torch.jit.script
def fn():
return math.pi
- Атрибуты ScriptModule можно пометить как константные, добавив аннотацию
Final[T]
import torch
import torch.nn as nn
class Foo(nn.Module):
# `Final` from the `typing_extensions` module can also be used
a : torch.jit.Final[int]
def __init__(self):
super(Foo, self).__init__()
self.a = 1 + 4
def forward(self, input):
return self.a + input
f = torch.jit.script(Foo())
Поддерживаемые типы констант Python:
intfloatbooltorch.devicetorch.layouttorch.dtype- кортежи, содержащие поддерживаемые типы
-
torch.nn.ModuleListкоторые могут использоваться в цикле TorchScript
Атрибуты модуля
Оборачивание torch.nn.Parameter и register_buffer могут быть использованы для присвоения тензоров модулю. Другие значения, присваиваемые скомпилированному модулю, будут добавлены в скомпилированный модуль, если их типы можно вывести. Все типы, доступные в TorchScript, могут использоваться в качестве атрибутов модуля. Атрибуты тензоров семантически аналогичны буферам. Тип пустых списков и словарей, а также None значений нельзя вывести, и их необходимо указать с помощью аннотаций в стиле PEP 526. Если тип нельзя вывести и он не аннотирован явно, он не будет добавлен в качестве атрибута в результирующий ScriptModule.
Пример:
from typing import List, Dict
class Foo(nn.Module):
# `words` is initialized as an empty list, so its type must be specified
words: List[str]
# The type could potentially be inferred if `a_dict` (below) was not
# empty, but this annotation ensures `some_dict` will be made into the
# proper type
some_dict: Dict[str, int]
def __init__(self, a_dict):
super(Foo, self).__init__()
self.words = []
self.some_dict = a_dict
# `int`s can be inferred
self.my_int = 10
def forward(self, input):
# type: (str) -> int
self.words.append(input)
return self.some_dict[input] + self.my_int
f = torch.jit.script(Foo({'hi': 2}))
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/jit_language_reference.html