TorchScript
- Создание кода TorchScript
- Смешивание трассировки и скриптирования
- Язык TorchScript
- Часто задаваемые вопросы
- Известные проблемы
TorchScript — это способ создания сериализуемых и оптимизируемых моделей из кода PyTorch. Любой скрипт TorchScript может быть сохранен из процесса Python и загружен в процесс, где нет зависимости от Python.
Мы предоставляем инструменты для поэтапного перехода модели от чистого кода Python к программе TorchScript, которая может выполняться независимо от Python, например, в автономной программе на C++. Это позволяет обучать модели в PyTorch с помощью привычных инструментов Python, а затем экспортировать модель через TorchScript в производственную среду, где программы Python могут быть невыгодны по производительности и причинам многопоточности.
Для знакомства с TorchScript см. учебное пособие Введение в TorchScript.
Для примера преобразования модели PyTorch в TorchScript и её выполнения в C++ см. учебное пособие Загрузка модели PyTorch в C++.
Создание кода TorchScript
script
| Скриптирование функции или |
trace
| Трассировка функции и возврат исполняемого или |
script_if_tracing
| Компилирует |
trace_module
| Трассировка модуля и возврат исполняемого |
fork
| Создаёт асинхронную задачу, выполняющую |
wait
| Заставляет завершить асинхронную задачу, возвращая результат задачи. |
ScriptModule
| Обёртка вокруг C++ |
ScriptFunction
| Функционально эквивалентно |
freeze
| Замораживание |
optimize_for_inference
| Выполняет набор оптимизационных проходов для оптимизации модели для целей вывода. |
enable_onednn_fusion
| Включает или отключает onednn JIT слияние, в зависимости от параметра |
onednn_fusion_enabled
| Возвращает, включено ли onednn JIT слияние |
set_fusion_strategy
| Устанавливает тип и количество специализаций, которые могут произойти во время слияния. |
strict_fusion
| Этот класс генерирует ошибку, если не все узлы не были объединены при выводе или символично дифференцированы при обучении. |
save
| Сохранить автономную версию этого модуля для использования в отдельном процессе. |
load
| Загрузить |
ignore
| Этот декоратор указывает компилятору, что функция или метод должны быть проигнорированы и оставлены как функция Python. |
unused
| Этот декоратор указывает компилятору, что функция или метод должны быть проигнорированы и заменены на выброс исключения. |
isinstance
| Эта функция обеспечивает уточнение типов контейнеров в TorchScript. |
Attribute
| Этот метод является проходной функцией, которая возвращает |
annotate
| Этот метод является проходной функцией, которая возвращает |
Смешивание трассировки и скриптинга
Во многих случаях трассировка или скриптинг являются более простым подходом для преобразования модели в TorchScript. Трассировка и скриптинг могут быть объединены для удовлетворения конкретных требований части модели.
Скриптированные функции могут вызывать отслеженные функции. Это особенно полезно, когда вам нужно использовать управляющую последовательность вокруг простой модели прямого распространения. Например, поиск по путям (beam search) в модели «последовательность к последовательности» обычно записывается в скрипте, но может вызывать модуль кодера, сгенерированный с помощью трассировки.
Пример (вызов отслеженной функции в скрипте):
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.Module, где его можно использовать для генерации подмодуля с помощью трассировки, который можно вызывать из методов скриптового модуля.
Пример (использование отслеженного модуля):
import torch
import torchvision
class MyScriptModule(torch.nn.Module):
def __init__(self):
super(MyScriptModule, self).__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, и поэтому при другой трассировке с другой формой граф отличается.
В этом случае управляющую последовательность, зависящую от данных, можно захватить с помощью 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(Model, self).__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(Model, self).__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(MyModule, self).__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(MyModule, self).__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(MyModule, self).__init__()
self.my_constant = 2
def forward(self):
pass
m = MyModule()
Новый API:
try:
from typing_extensions import Final
except:
# If you don't have `typing_extensions` installed, you can use a
# polyfill from `torch.jit`.
from torch.jit import Final
class MyModule(torch.nn.Module):
my_constant: Final[int]
def __init__(self):
super(MyModule, self).__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 для получения более подробной информации об использовании и отладке.
Ссылки
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/jit.html