Справочник по языку 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, либо имеет тип |
| Словарь с ключом типа |
| |
| |
| Тип кортежа |
| Один из подтипов |
В отличие от 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().__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().__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, как в TorchScript, так и в Python:
# 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, как в TorchScript, так и в Python. Тип значений перечисления должен быть 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 и torch.nn.ModuleDict.
Выражения
Поддерживаются следующие выражения 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().__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
В дополнение к булевым значениям, в условных операторах также можно использовать числа с плавающей точкой, целые числа и тензоры, которые неявно будут преобразованы в булевые значения.
Циклы 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().__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().__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
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().__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().__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/2.1/jit_language_reference.html