Spec-Zone.ru › PyTorch 2

TorchScript

  • Справочник по языку TorchScript
  • Создание кода TorchScript
  • Смешение трассировки и скриптирования
  • Язык TorchScript
  • Встроенные функции и модули

    • Функции и модули PyTorch
    • Функции и модули Python
    • Сравнение со справочником по языку Python
  • Отладка

    • Отключение JIT для отладки
    • Просмотр кода
    • Интерпретация графиков
    • Трассировщик
  • Часто задаваемые вопросы
  • Известные проблемы
  • Приложение

    • Миграция на API рекурсивного скриптирования PyTorch 1.2
    • Фьюжн-бэкэнды
    • Ссылки

TorchScript — это способ создания сериализуемых и оптимизируемых моделей из кода PyTorch. Любой скрипт TorchScript можно сохранить из процесса Python и загрузить в процессе, где нет зависимости от Python.

Мы предоставляем инструменты для поэтапного перехода модели от чистого кода Python к коду TorchScript, который может выполняться независимо от Python, например, в автономной C++ программе. Это позволяет обучать модели в PyTorch с помощью знакомых инструментов в Python, а затем экспортировать модель через TorchScript в производственную среду, где программы Python могут быть невыгодными по соображениям производительности и многопоточности.

Для ознакомления с TorchScript, см. введение в TorchScript в руководстве.

Для примера преобразования модели PyTorch в TorchScript и её выполнения в C++, см. руководство Загрузка модели PyTorch в C++.

Создание кода TorchScript

script

Скриптирование функции или nn.Module проанализирует исходный код, скомпилирует его как код TorchScript с помощью компилятора TorchScript и вернёт ScriptModule или ScriptFunction.

trace

Трассировка функции и возврат исполняемого или ScriptFunction, который будет оптимизирован с помощью JIT-компиляции.

script_if_tracing

Компилирует fn при первом вызове во время трассировки.

trace_module

Трассировка модуля и возврат исполняемого ScriptModule, который будет оптимизирован с помощью JIT-компиляции.

fork

Создаёт асинхронную задачу, выполняющую func и ссылку на значение результата этого выполнения.

wait

Вынуждает завершение асинхронной задачи torch.jit.Future[T], возвращая результат задачи.

ScriptModule

Обёртка вокруг C++ torch::jit::Module.

ScriptFunction

Функционально эквивалентно ScriptModule, но представляет отдельную функцию и не имеет атрибутов или параметров.

freeze

Замораживание ScriptModule клонирует его и пытается встроить подмодули, параметры и атрибуты клонированного модуля как константы в графе TorchScript IR.

optimize_for_inference

Выполняет набор оптимизационных проходов для оптимизации модели для целей вывода.

enable_onednn_fusion

Включает или выключает onednn JIT-слияние на основе параметра enabled.

onednn_fusion_enabled

Возвращает, включено ли onednn JIT-слияние.

set_fusion_strategy

Устанавливает тип и количество специализаций, которые могут произойти во время слияния.

strict_fusion

Этот класс выдаёт ошибку, если не все узлы были объединены во время вывода или символично дифференцированы во время обучения.

save

Сохранение автономной версии данного модуля для использования в отдельном процессе.

load

Загрузка ScriptModule или ScriptFunction, ранее сохранённых с помощью torch.jit.save

ignore

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

unused

Этот декоратор указывает компилятору, что функция или метод должны быть проигнорированы и заменены на повышение исключения.

isinstance

Эта функция обеспечивает уточнение типа контейнера в TorchScript.

Attribute

Этот метод — функция-проход, которая возвращает value, в основном используется для указания компилятору TorchScript, что левая часть выражения — экземпляр атрибута класса с типом type.

annotate

Этот метод — функция-проход, которая возвращает the_value, используется для указания компилятору TorchScript типа the_value.

Смешение трассировки и скриптирования

Во многих случаях трассировка или скриптирование — это более простой подход для преобразования модели в TorchScript. Трассировка и скриптирование могут быть объединены, чтобы соответствовать конкретным требованиям части модели.

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

Пример (вызов отслеженной функции в скрипте):

import torch

def foo(x, y):
    return 2 * x + y

traced_foo = torch.jit.trace(foo, (torch.rand(3), torch.rand(3)))

@torch.jit.script
def bar(x):
    return traced_foo(x, x)

Отслеженные функции могут вызывать скриптированные функции. Это полезно, когда небольшая часть модели требует некоторой обработки потока управления, даже если большая часть модели представляет собой просто сеть прямой передачи данных. Обработка потока управления внутри скриптированной функции, вызываемой отслеженной функцией, сохраняется корректно.

Пример (вызов скриптированной функции в отслеченной функции):

import torch

@torch.jit.script
def foo(x, y):
    if x.max() > y.max():
        r = x
    else:
        r = y
    return r


def bar(x, y, z):
    return foo(x, y) + z

traced_bar = torch.jit.trace(bar, (torch.rand(3), torch.rand(3), torch.rand(3)))

Это составление также работает для nn.Modules, где оно может использоваться для генерации подмодуля с использованием трассировки, который может вызываться из методов модуля сценария.

Пример (с использованием отслеживаемого модуля):

import torch
import torchvision

class MyScriptModule(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.means = torch.nn.Parameter(torch.tensor([103.939, 116.779, 123.68])
                                        .resize_(1, 3, 1, 1))
        self.resnet = torch.jit.trace(torchvision.models.resnet18(),
                                      torch.rand(1, 3, 224, 224))

    def forward(self, input):
        return self.resnet(input - self.means)

my_script_module = torch.jit.script(MyScriptModule())

Язык TorchScript

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

Встроенные функции и модули

TorchScript поддерживает использование большинства функций PyTorch и многих встроенных функций Python. Полный список поддерживаемых функций см. в Встроенных функциях TorchScript.

Функции и модули PyTorch

TorchScript поддерживает подмножество функций тензора и нейронной сети, предоставляемых PyTorch. Большинство методов тензора, а также функции в пространстве имен torch, все функции в torch.nn.functional и большинство модулей из torch.nn поддерживаются в TorchScript.

Список неподдерживаемых функций и модулей PyTorch см. в Неподдерживаемые конструкции PyTorch в TorchScript.

Функции и модули Python

Многие встроенные функции Python поддерживаются в TorchScript. Модуль math также поддерживается (подробности см. в Модуле math), но другие модули Python (встроенные или сторонние) не поддерживаются.

Сравнение со справочником по языку Python

Полный список поддерживаемых функций Python см. в Охвате справочника по языку Python.

Отладка

Отключение JIT для отладки

PYTORCH_JIT

Установив переменную среды PYTORCH_JIT=0, вы отключите все аннотации сценария и трассировки. Если в одной из ваших моделей TorchScript есть ошибка, сложная для отладки, вы можете использовать этот флаг для принудительного запуска всего с использованием обычного Python. Поскольку TorchScript (скрипты и трассировка) отключены с этим флагом, вы можете использовать такие инструменты, как pdb, для отладки кода модели. Например:

@torch.jit.script
def scripted_fn(x : torch.Tensor):
    for i in range(12):
        x = x + x
    return x

def fn(x):
    x = torch.neg(x)
    import pdb; pdb.set_trace()
    return scripted_fn(x)

traced_fn = torch.jit.trace(fn, (torch.rand(4, 5),))
traced_fn(torch.rand(3, 4))

Отладка этого скрипта с помощью pdb работает, за исключением случаев, когда мы вызываем функцию @torch.jit.script. Мы можем глобально отключить JIT, чтобы мы могли вызывать функцию @torch.jit.script как обычную функцию Python, а не компилировать её. Если вышеуказанный скрипт называется disable_jit_example.py, мы можем вызвать его так:

$ PYTORCH_JIT=0 python disable_jit_example.py

и мы сможем войти в функцию @torch.jit.script как в обычную функцию Python. Для отключения компилятора TorchScript для конкретной функции см. @torch.jit.ignore.

Просмотр кода

TorchScript предоставляет красивое представление кода для всех экземпляров ScriptModule. Это красивое представление показывает интерпретацию кода метода скрипта как корректного синтаксиса Python. Например:

@torch.jit.script
def foo(len):
    # type: (int) -> torch.Tensor
    rv = torch.zeros(3, 4)
    for i in range(len):
        if i < 10:
            rv = rv - 1.0
        else:
            rv = rv + 1.0
    return rv

print(foo.code)

У ScriptModule с одним методом forward будет атрибут code, который вы можете использовать для просмотра кода ScriptModule. Если ScriptModule имеет более одного метода, вам нужно будет получить доступ к .code к самому методу, а не к модулю. Мы можем посмотреть код метода с именем foo в ScriptModule, получив доступ к .foo.code. Приведённый выше пример выводит:

def foo(len: int) -> Tensor:
    rv = torch.zeros([3, 4], dtype=None, layout=None, device=None, pin_memory=None)
    rv0 = rv
    for i in range(len):
        if torch.lt(i, 10):
            rv1 = torch.sub(rv0, 1., 1)
        else:
            rv1 = torch.add(rv0, 1., 1)
        rv0 = rv1
    return rv0

Это компиляция TorchScript кода для метода forward. Вы можете использовать это, чтобы убедиться, что TorchScript (трассировка или скрипты) правильно захватил ваш код модели.

Интерпретация графиков

TorchScript также имеет представление на уровне ниже, чем красивое представление кода, в виде графиков IR.

TorchScript использует статическое представление с единым назначением (SSA) промежуточного представления (IR) для представления вычислений. Инструкции в этом формате состоят из операторов ATen (C++ бэкэнд PyTorch) и других примитивных операторов, включая операторы управления потоком для циклов и условных выражений. Пример:

@torch.jit.script
def foo(len):
    # type: (int) -> torch.Tensor
    rv = torch.zeros(3, 4)
    for i in range(len):
        if i < 10:
            rv = rv - 1.0
        else:
            rv = rv + 1.0
    return rv

print(foo.graph)

graph следует тем же правилам, что описаны в разделе Просмотр кода относительно поиска метода forward.

Приведенный выше пример скрипта генерирует график:

graph(%len.1 : int):
  %24 : int = prim::Constant[value=1]()
  %17 : bool = prim::Constant[value=1]() # test.py:10:5
  %12 : bool? = prim::Constant()
  %10 : Device? = prim::Constant()
  %6 : int? = prim::Constant()
  %1 : int = prim::Constant[value=3]() # test.py:9:22
  %2 : int = prim::Constant[value=4]() # test.py:9:25
  %20 : int = prim::Constant[value=10]() # test.py:11:16
  %23 : float = prim::Constant[value=1]() # test.py:12:23
  %4 : int[] = prim::ListConstruct(%1, %2)
  %rv.1 : Tensor = aten::zeros(%4, %6, %6, %10, %12) # test.py:9:10
  %rv : Tensor = prim::Loop(%len.1, %17, %rv.1) # test.py:10:5
    block0(%i.1 : int, %rv.14 : Tensor):
      %21 : bool = aten::lt(%i.1, %20) # test.py:11:12
      %rv.13 : Tensor = prim::If(%21) # test.py:11:9
        block0():
          %rv.3 : Tensor = aten::sub(%rv.14, %23, %24) # test.py:12:18
          -> (%rv.3)
        block1():
          %rv.6 : Tensor = aten::add(%rv.14, %23, %24) # test.py:14:18
          -> (%rv.6)
      -> (%17, %rv.13)
  return (%rv)

Рассмотрим инструкцию %rv.1 : Tensor = aten::zeros(%4, %6, %6, %10, %12) # test.py:9:10 как пример.

  • %rv.1 : Tensor означает, что мы присваиваем выходное значение (уникальному) значению с именем rv.1, это значение типа Tensor, и мы не знаем его конкретную форму.
  • aten::zeros — оператор (эквивалентный torch.zeros), а список входных данных (%4, %6, %6, %10, %12) указывает, какие значения в области видимости должны передаваться в качестве входных данных. Схему для встроенных функций, таких как aten::zeros, см. в разделе Встроенные функции.
  • # test.py:9:10 — расположение в исходном файле, которое сгенерировало эту инструкцию. В данном случае это файл с именем test.py, строка 9, символ 10.

Обратите внимание, что операторы также могут иметь ассоциированные blocks, а именно операторы prim::Loop и prim::If. В выводе графика эти операторы отформатированы, чтобы отразить их эквивалентные формы исходного кода, что облегчит отладку.

Графики можно просмотреть, как показано выше, чтобы подтвердить, что вычисление, описанное ScriptModule, верно как в автоматическом, так и в ручном режиме, как описано ниже.

Трассировщик

Крайние случаи трассировки

Существуют некоторые крайние случаи, в которых трассировка данной функции/модуля Python не будет отражать исходный код. Эти случаи могут включать:

  • Трассировка потока управления, зависящего от входных данных (например, форм тензоров)
  • Трассировка операций на месте с представлениями тензоров (например, индексация в левой части присваивания)

Обратите внимание, что эти случаи могут быть отслеживаемы в будущем.

Автоматическая проверка трассировки

Один из способов автоматически обнаружить многие ошибки в трассировках — использовать check_inputs в API torch.jit.trace(). check_inputs принимает список кортежей входных данных, которые будут использоваться для повторной трассировки вычислений и проверки результатов. Например:

def loop_in_traced_fn(x):
    result = x[0]
    for i in range(x.size(0)):
        result = result * x[i]
    return result

inputs = (torch.rand(3, 4, 5),)
check_inputs = [(torch.rand(4, 5, 6),), (torch.rand(2, 3, 4),)]

traced = torch.jit.trace(loop_in_traced_fn, inputs, check_inputs=check_inputs)

Это даёт следующую диагностическую информацию:

ERROR: Graphs differed across invocations!
Graph diff:

            graph(%x : Tensor) {
            %1 : int = prim::Constant[value=0]()
            %2 : int = prim::Constant[value=0]()
            %result.1 : Tensor = aten::select(%x, %1, %2)
            %4 : int = prim::Constant[value=0]()
            %5 : int = prim::Constant[value=0]()
            %6 : Tensor = aten::select(%x, %4, %5)
            %result.2 : Tensor = aten::mul(%result.1, %6)
            %8 : int = prim::Constant[value=0]()
            %9 : int = prim::Constant[value=1]()
            %10 : Tensor = aten::select(%x, %8, %9)
        -   %result : Tensor = aten::mul(%result.2, %10)
        +   %result.3 : Tensor = aten::mul(%result.2, %10)
        ?          ++
            %12 : int = prim::Constant[value=0]()
            %13 : int = prim::Constant[value=2]()
            %14 : Tensor = aten::select(%x, %12, %13)
        +   %result : Tensor = aten::mul(%result.3, %14)
        +   %16 : int = prim::Constant[value=0]()
        +   %17 : int = prim::Constant[value=3]()
        +   %18 : Tensor = aten::select(%x, %16, %17)
        -   %15 : Tensor = aten::mul(%result, %14)
        ?     ^                                 ^
        +   %19 : Tensor = aten::mul(%result, %18)
        ?     ^                                 ^
        -   return (%15);
        ?             ^
        +   return (%19);
        ?             ^
            }

Это сообщение указывает, что вычисление отличалось между первым трассированием и трассированием с check_inputs. Действительно, цикл внутри тела loop_in_traced_fn зависит от формы входных данных x, и поэтому, когда мы пытаемся ещё одно x с другой формой, трассировка отличается.

В этом случае потоки управления, зависящие от данных, можно захватить с помощью torch.jit.script() вместо этого:

def fn(x):
    result = x[0]
    for i in range(x.size(0)):
        result = result * x[i]
    return result

inputs = (torch.rand(3, 4, 5),)
check_inputs = [(torch.rand(4, 5, 6),), (torch.rand(2, 3, 4),)]

scripted_fn = torch.jit.script(fn)
print(scripted_fn.graph)
#print(str(scripted_fn.graph).strip())

for input_tuple in [inputs] + check_inputs:
    torch.testing.assert_close(fn(*input_tuple), scripted_fn(*input_tuple))

Что даёт:

graph(%x : Tensor) {
    %5 : bool = prim::Constant[value=1]()
    %1 : int = prim::Constant[value=0]()
    %result.1 : Tensor = aten::select(%x, %1, %1)
    %4 : int = aten::size(%x, %1)
    %result : Tensor = prim::Loop(%4, %5, %result.1)
    block0(%i : int, %7 : Tensor) {
        %10 : Tensor = aten::select(%x, %1, %i)
        %result.2 : Tensor = aten::mul(%7, %10)
        -> (%5, %result.2)
    }
    return (%result);
}

Предупреждения трассировщика

Трассировщик генерирует предупреждения для нескольких проблемных паттернов в прослеживаемых вычислениях. Например, рассмотрим трассировку функции, содержащей присвоение на месте для среза (представления) тензора:

def fill_row_zero(x):
    x[0] = torch.rand(*x.shape[1:2])
    return x

traced = torch.jit.trace(fill_row_zero, (torch.rand(3, 4),))
print(traced.graph)

Генерирует несколько предупреждений и график, который просто возвращает входные данные:

fill_row_zero.py:4: TracerWarning: There are 2 live references to the data region being modified when tracing in-place operator copy_ (possibly due to an assignment). This might cause the trace to be incorrect, because all other views that also reference this data will not reflect this change in the trace! On the other hand, if all other views use the same memory chunk, but are disjoint (e.g. are outputs of torch.split), this might still be safe.
    x[0] = torch.rand(*x.shape[1:2])
fill_row_zero.py:6: TracerWarning: Output nr 1. of the traced function does not match the corresponding output of the Python function. Detailed error:
Not within tolerance rtol=1e-05 atol=1e-05 at input[0, 1] (0.09115803241729736 vs. 0.6782537698745728) and 3 other locations (33.00%)
    traced = torch.jit.trace(fill_row_zero, (torch.rand(3, 4),))
graph(%0 : Float(3, 4)) {
    return (%0);
}

Можно исправить это, изменив код, чтобы не использовать обновление на месте, а вместо этого создавать результирующий тензор вне места с помощью torch.cat:

def fill_row_zero(x):
    x = torch.cat((torch.rand(1, *x.shape[1:2]), x[1:2]), dim=0)
    return x

traced = torch.jit.trace(fill_row_zero, (torch.rand(3, 4),))
print(traced.graph)

Часто задаваемые вопросы

В: Я хочу обучать модель на GPU и делать инференс на CPU. Какие лучшие практики?

Сначала преобразуйте свою модель с GPU на CPU и сохраните её, как показано ниже:

cpu_model = gpu_model.cpu()
sample_input_cpu = sample_input_gpu.cpu()
traced_cpu = torch.jit.trace(cpu_model, sample_input_cpu)
torch.jit.save(traced_cpu, "cpu.pt")

traced_gpu = torch.jit.trace(gpu_model, sample_input_gpu)
torch.jit.save(traced_gpu, "gpu.pt")

# ... later, when using the model:

if use_gpu:
  model = torch.jit.load("gpu.pt")
else:
  model = torch.jit.load("cpu.pt")

model(input)

Это рекомендуется, потому что трассировщик может наблюдать создание тензора на определённом устройстве, поэтому преобразование уже загруженной модели может иметь непредвиденные последствия. Преобразование модели до её сохранения гарантирует, что трассировщик имеет правильную информацию об устройстве.

В: Как сохранить атрибуты в ScriptModule?

Допустим, у нас есть модель, как:

import torch

class Model(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.x = 2

    def forward(self):
        return self.x

m = torch.jit.script(Model())

Если Model будет создан, это приведёт к ошибке компиляции, поскольку компилятор не знает о x. Существует 4 способа сообщить компилятору об атрибутах в ScriptModule:

1. nn.Parameter - Значения, заключённые в nn.Parameter, будут работать так же, как и в nn.Module.

2. register_buffer - Значения, заключённые в register_buffer, будут работать так же, как и в nn.Module. Это эквивалентно атрибуту (см. 4) типа Tensor.

3. Константы - Аннотирование члена класса как Final (или добавление его в список, называемый __constants__ на уровне определения класса) помечает содержащиеся имена как константы. Константы сохраняются непосредственно в коде модели. Подробнее см. builtin-constants.

4. Атрибуты - Значения, являющиеся supported type, могут быть добавлены в качестве изменяемых атрибутов. Большинство типов могут быть выведены, но некоторые могут потребовать явного указания, см. module attributes для получения подробностей.

В: Я хотел бы отслеживать метод модуля, но продолжаю получать эту ошибку:

RuntimeError: Cannot insert a Tensor that requires grad as a constant. Consider making it a parameter or input, or detaching the gradient

Эта ошибка обычно означает, что метод, который вы отслеживаете, использует параметры модуля, и вы передаёте метод модуля вместо экземпляра модуля (например, my_module_instance.forward вместо my_module_instance).

  • Вызов trace с методом модуля фиксирует параметры модуля (которые могут потребовать градиенты) как константы.
  • С другой стороны, вызов trace с экземпляром модуля (например, my_module) создаёт новый модуль и правильно копирует параметры в новый модуль, поэтому они могут накапливать градиенты, если это необходимо.

Для отслеживания конкретного метода в модуле, см. torch.jit.trace_module

Известные проблемы

Если вы используете Sequential с TorchScript, входные данные некоторых Sequential подмодулей могут быть ложно выведены как Tensor, даже если они аннотированы иначе. Стандартное решение - создать подкласс nn.Sequential и повторно объявить forward с правильным типом входных данных.

Приложение

Миграция к API рекурсивного скриптинга PyTorch 1.2

Этот раздел подробно описывает изменения в TorchScript в PyTorch 1.2. Если вы новичок в TorchScript, можете пропустить этот раздел. В API TorchScript с PyTorch 1.2 есть два основных изменения.

1. torch.jit.script теперь будет пытаться рекурсивно компилировать функции, методы и классы, которые он встречает. После вызова torch.jit.script, компиляция становится «выключенной по умолчанию», а не «включенной по умолчанию».

2. torch.jit.script(nn_module_instance) теперь является предпочтительным способом создания ScriptModule вместо наследования от torch.jit.ScriptModule. Эти изменения объединяются, чтобы предоставить более простой и удобный в использовании API для преобразования ваших nn.Module в ScriptModule, готовые к оптимизации и выполнению в среде без Python.

Новый способ использования:

import torch
import torch.nn as nn
import torch.nn.functional as F

class Model(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(1, 20, 5)
        self.conv2 = nn.Conv2d(20, 20, 5)

    def forward(self, x):
        x = F.relu(self.conv1(x))
        return F.relu(self.conv2(x))

my_model = Model()
my_scripted_model = torch.jit.script(my_model)
  • Модуль forward компилируется по умолчанию. Методы, вызываемые из forward, компилируются лениво в порядке их использования в forward.
  • Для компиляции метода, отличного от forward, который не вызывается из forward, добавьте @torch.jit.export.
  • Для остановки компиляции метода добавьте @torch.jit.ignore или @torch.jit.unused. @ignore оставляет метод как вызов на Python, а @unused заменяет его исключением. @ignored не может быть экспортирован; @unused может.
  • Большинство типов атрибутов могут быть выведены, поэтому torch.jit.Attribute не обязательно. Для пустых контейнерных типов укажите их типы с помощью аннотаций классов в стиле PEP 526.
  • Константы можно пометить с помощью аннотации класса Final вместо добавления имени члена в __constants__.
  • Можно использовать подсказки типов Python 3 вместо torch.jit.annotate.
В результате этих изменений следующие элементы считаются устаревшими и не должны появляться в новом коде:
  • Декоратор @torch.jit.script_method
  • Классы, наследующие от torch.jit.ScriptModule
  • Обёртный класс torch.jit.Attribute
  • Массив __constants__
  • Функция torch.jit.annotate

Модули

Предупреждение

Поведение аннотации @torch.jit.ignore меняется в PyTorch 1.2. До PyTorch 1.2 декоратор @ignore использовался для того, чтобы сделать функцию или метод вызываемым из экспортируемого кода. Для восстановления этого функционала используйте @torch.jit.unused(). @torch.jit.ignore теперь эквивалентно @torch.jit.ignore(drop=False) . См. @torch.jit.ignore и @torch.jit.unused для подробностей.

При передаче в функцию torch.jit.script, данные torch.nn.Module копируются в ScriptModule, и компилятор TorchScript компилирует модуль. Модуль forward компилируется по умолчанию. Методы, вызываемые из forward, компилируются лениво в порядке их использования в forward, а также любые методы @torch.jit.export.

torch.jit.export(fn) [source]

Этот декоратор указывает, что метод в nn.Module используется в качестве точки входа в ScriptModule и должен быть скомпилирован.

forward неявно предполагается точкой входа, поэтому этот декоратор для него не нужен. Функции и методы, вызываемые из forward компилируются по мере их обнаружения компилятором, поэтому они также не нуждаются в этом декораторе.

Пример (использование @torch.jit.export для метода):

import torch
import torch.nn as nn

class MyModule(nn.Module):
    def implicitly_compiled_method(self, x):
        return x + 99

    # `forward` is implicitly decorated with `@torch.jit.export`,
    # so adding it here would have no effect
    def forward(self, x):
        return x + 10

    @torch.jit.export
    def another_forward(self, x):
        # When the compiler sees this call, it will compile
        # `implicitly_compiled_method`
        return self.implicitly_compiled_method(x)

    def unused_method(self, x):
        return x - 20

# `m` will contain compiled methods:
#     `forward`
#     `another_forward`
#     `implicitly_compiled_method`
# `unused_method` will not be compiled since it was not called from
# any compiled methods and wasn't decorated with `@torch.jit.export`
m = torch.jit.script(MyModule())

Функции

Функции не сильно меняются, их можно украсить @torch.jit.ignore или torch.jit.unused при необходимости.

# Same behavior as pre-PyTorch 1.2
@torch.jit.script
def some_fn():
    return 2

# Marks a function as ignored, if nothing
# ever calls it then this has no effect
@torch.jit.ignore
def some_fn2():
    return 2

# As with ignore, if nothing calls it then it has no effect.
# If it is called in script it is replaced with an exception.
@torch.jit.unused
def some_fn3():
  import pdb; pdb.set_trace()
  return 4

# Doesn't do anything, this function is already
# the main entry point
@torch.jit.export
def some_fn4():
    return 2

Классы TorchScript

Предупреждение

Поддержка классов TorchScript находится в стадии разработки. В настоящее время она лучше всего подходит для простых типов, похожих на записи (например, NamedTuple с присоединёнными методами).

Всё в пользовательском классе TorchScript экспортируется по умолчанию; функции можно украсить @torch.jit.ignore, если нужно.

Атрибуты

Компилятор TorchScript должен знать типы module attributes. Большинство типов могут быть выведены из значения члена. Пустые списки и словари не могут иметь свои типы выведены и должны иметь их типы аннотированы с помощью аннотаций классов в стиле PEP 526. Если тип не может быть выведен и не аннотирован явно, он не будет добавлен в качестве атрибута в результирующий ScriptModule.

Старый API:

from typing import Dict
import torch

class MyModule(torch.jit.ScriptModule):
    def __init__(self):
        super().__init__()
        self.my_dict = torch.jit.Attribute({}, Dict[str, int])
        self.my_int = torch.jit.Attribute(20, int)

m = MyModule()

Новый API:

from typing import Dict

class MyModule(torch.nn.Module):
    my_dict: Dict[str, int]

    def __init__(self):
        super().__init__()
        # This type cannot be inferred and must be specified
        self.my_dict = {}

        # The attribute type here is inferred to be `int`
        self.my_int = 20

    def forward(self):
        pass

m = torch.jit.script(MyModule())

Константы

Конструктор типа Final может использоваться для пометки членов как constant. Если члены не помечены как константы, они будут скопированы в результирующий ScriptModule как атрибут. Использование Final открывает возможности для оптимизации, если значение известно как фиксированное, и обеспечивает дополнительную типобезопасность.

Старый API:

class MyModule(torch.jit.ScriptModule):
    __constants__ = ['my_constant']

    def __init__(self):
        super().__init__()
        self.my_constant = 2

    def forward(self):
        pass
m = MyModule()

Новый API:

from typing import Final

class MyModule(torch.nn.Module):

    my_constant: Final[int]

    def __init__(self):
        super().__init__()
        self.my_constant = 2

    def forward(self):
        pass

m = torch.jit.script(MyModule())

Переменные

Контейнеры предполагаются с типом Tensor и не являются необязательными (см. Default Types для получения дополнительной информации). Раньше torch.jit.annotate использовалось для указания компилятору TorchScript, какой тип должен быть. Теперь поддерживаются подсказки типов в стиле Python 3.

import torch
from typing import Dict, Optional

@torch.jit.script
def make_dict(flag: bool):
    x: Dict[str, int] = {}
    x['hi'] = 2
    b: Optional[int] = None
    if flag:
        b = 2
    return x, b

Бэкенды слияния

Существует несколько бэкэндов слияния для оптимизации выполнения TorchScript. По умолчанию fuser на процессорах - NNC, который может выполнять слияния для процессоров и графических процессоров. По умолчанию fuser на графических процессорах - NVFuser, который поддерживает более широкий спектр операторов и продемонстрировал ядра с улучшенной пропускной способностью. См. документацию NVFuser для получения более подробной информации об использовании и отладке.

Ссылки

  • Справочная информация по языку Python
  • Неподдерживаемые конструкции PyTorch в TorchScript

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/jit.html

Spec-Zone.ru

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