Spec-Zone.ru › PyTorch 2
  • Типы
  • Выражения
  • Операторы
  • Разрешение переменных
  • Использование значений Python

Справочник по языку 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 поддерживает:

Тип

Описание

Tensor

Тензор PyTorch любого типа данных, размерности или бэкенда

Tuple[T0, T1, ..., TN]

Кортеж, содержащий подтипы T0, T1, и т. д. (например, Tuple[Tensor, Tensor])

bool

Булево значение

int

Скалярное целое число

float

Скалярное число с плавающей точкой

str

Строка

List[T]

Список, все члены которого имеют тип T

Optional[T]

Значение, которое либо None, либо имеет тип T

Dict[K, V]

Словарь с ключом типа K и значением типа V. Допустимы только ключи типов str, int, и float.

T

Класс TorchScript

E

Перечисление TorchScript

NamedTuple[T0, T1, ...]

Тип кортежа collections.namedtuple

Union[T0, T1, ...]

Один из подтипов T0, T1, и т. д.

В отличие от 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.

Элемент

Описание

typing.Any

typing.Any в настоящее время находится в разработке, но еще не выпущена

typing.NoReturn

Не реализовано

typing.Sequence

Не реализовано

typing.Callable

Не реализовано

typing.Literal

Не реализовано

typing.ClassVar

Не реализовано

typing.Final

Поддерживается для аннотаций атрибутов модуля, но не для функций

typing.AnyStr

TorchScript не поддерживает bytes, поэтому этот тип не используется

typing.overload

typing.overload в настоящее время находится в разработке, но еще не выпущена

Псевдонимы типов

Не реализовано

Номинальное против структурного подтипирования

Номинальная типизация находится в разработке, но структурная — нет

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)
Возвращаемый тип

bool

torch.jit.is_tracing() [source]

Возвращает True при трассировке (если функция вызывается во время трассировки кода с torch.jit.trace ) и False в противном случае.

Поиск атрибутов в модулях Python

TorchScript может искать атрибуты в модулях. Builtin functions как torch.add доступны таким образом. Это позволяет TorchScript вызывать функции, определенные в других модулях.

Постоянные значения, определенные в Python

TorchScript также предоставляет способ использования констант, определенных в Python. Их можно использовать для жесткого кодирования гиперпараметров в функцию или для определения универсальных констант. Существует два способа указать, что значение Python должно обрабатываться как константа.

  1. Значения, ищущиеся как атрибуты модуля, предполагаются постоянными:
import math
import torch

@torch.jit.script
def fn():
    return math.pi
  1. Атрибуты 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

  • int
  • float
  • bool
  • torch.device
  • torch.layout
  • torch.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

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API