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. Трассировка и скриптирование могут быть объединены, чтобы соответствовать конкретным требованиям части модели.
Скриптированные функции могут вызывать отслеженные функции. Это особенно полезно, когда вам нужна обработка потока управления вокруг простой модели прямой передачи данных. Например, поиск по лучшим вариантам в модели последовательности к последовательности обычно записывается в скрипте, но может вызывать модуль кодировщика, сгенерированный с помощью трассировки.
Пример (вызов отслеженной функции в скрипте):
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 для получения более подробной информации об использовании и отладке.
Ссылки
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/jit.html