torch.jit.trace
-
torch.jit.trace(func, example_inputs, optimize=None, check_trace=True, check_inputs=None, check_tolerance=1e-05, strict=True, _force_outplace=False, _module_class=None, _compilation_unit=<torch.jit.CompilationUnit object>)[source] -
Отслеживает функцию и возвращает исполняемый или
ScriptFunctionобъект, который будет оптимизирован с помощью компиляции в момент использования. Отслеживание идеально подходит для кода, который работает только сTensorи списками, словарями и кортежамиTensor.Используя
torch.jit.traceиtorch.jit.trace_module, вы можете преобразовать существующий модуль или Python-функцию в TorchScriptScriptFunctionилиScriptModule. Вам необходимо предоставить примеры входных данных, и мы запустим функцию, записывая операции, выполняемые над всеми тензорами.- Результат записи автономной функции создаёт
ScriptFunction. - Результат записи
nn.Module.forwardилиnn.ModuleсоздаётScriptModule.
Этот модуль также содержит любые параметры, которые имел исходный модуль.
Предупреждение
Отслеживание корректно записывает только функции и модули, которые не зависят от данных (например, не содержат условных операторов на основе данных в тензорах) и не имеют каких-либо неотслеживаемых внешних зависимостей (например, не выполняют ввод/вывод или доступ к глобальным переменным). Отслеживание записывает только операции, выполняемые, когда заданная функция выполняется на заданных тензорах. Таким образом, возвращённый
ScriptModuleвсегда выполняет один и тот же отслеженный граф для любого входного набора.- Отслеживание не будет записывать никакого управления потоком, например, операторов if-else или циклов. Когда это управление потоком является постоянным в вашем модуле, это нормально и часто встраивает решения о управлении потоком. Но иногда управление потоком на самом деле является частью модели. Например, рекуррентная сеть является циклом по (возможно, динамической) длине последовательности входных данных.
- В возвращённом
ScriptModule, операции, которые имеют различное поведение в режимахtrainingиeval, всегда будут вести себя так, как если бы они были в том режиме, в котором они находились во время отслеживания, независимо от того, в каком режиме находитсяScriptModule.
В таких случаях отслеживание не подходит, и
scriptingявляется лучшим выбором. Если вы отслеживаете такие модели, вы можете получить неверные результаты при последующих вызовах модели. Трассер попытается выдать предупреждения, когда выполняет действия, которые могут привести к созданию неверного отслеживания.- Parameters:
-
-
func (callable or torch.nn.Module) – Python-функция или
torch.nn.Module, которая будет запущена сexample_inputs. Аргументыfuncи возвращаемые значения должны быть тензорами или (возможно, вложенными) кортежами, содержащими тензоры. Когда передаётся модульtorch.jit.trace, выполняется и отслеживается только методforward(см.torch.jit.traceдля подробностей). -
example_inputs (tuple or torch.Tensor) – Кортеж примеров входных данных, которые будут переданы функции во время отслеживания. Результат отслеживания может быть запущен с входными данными разных типов и форм, при условии, что отслеживаемые операции поддерживают эти типы и формы.
example_inputsтакже может быть единственным тензором, в этом случае он автоматически упаковывается в кортеж.
-
func (callable or torch.nn.Module) – Python-функция или
- Keyword Arguments:
-
-
check_trace (
bool, optional) – Проверить, дают ли одни и те же входные данные одинаковые выходные результаты в отслеженном коде. По умолчанию:True. Возможно, вам захочется отключить это, если, например, ваша сеть содержит не детерминированные операции или если вы уверены, что сеть верна, несмотря на сбой проверки. -
check_inputs (list of tuples, optional) – Список кортежей аргументов входных данных, которые должны быть использованы для проверки отслеживания по сравнению с ожидаемым результатом. Каждый кортеж эквивалентен набору аргументов входных данных, которые были бы указаны в
example_inputs. Для наилучших результатов передайте набор входных данных проверки, представляющий пространство форм и типов входных данных, которые вы ожидаете увидеть в сети. Если не указано, то исходныеexample_inputsиспользуются для проверки. - check_tolerance (float, optional) – Допуск сравнения с плавающей точкой, используемый в процедуре проверки. Это может быть использовано для ослабления строгости проверки в случае, если результаты расходятся численно по известной причине, такой как слияние операторов.
-
strict (
bool, optional) – Запустить трассер в строгом режиме или нет (по умолчанию:True). Отключайте только в случае, если хотите, чтобы трассер записывал ваши изменяемые типы контейнеров (в настоящее времяlist/dict) и уверены, что контейнер, который вы используете в своей задаче, представляет собой структуруconstantи не используется в качестве условий управления потоком (if, for).
-
check_trace (
- Returns:
-
Если
funcявляетсяnn.Moduleилиforwardизnn.Module,traceвозвращает объектScriptModuleс единственным методомforward, содержащим отслеженный код. ВозвращённыйScriptModuleбудет иметь тот же набор подмодулей и параметров, что и исходныйnn.Module. Еслиfuncявляется автономной функцией,traceвозвращаетScriptFunction.
Пример (отслеживание функции):
import torch def foo(x, y): return 2 * x + y # Run `foo` with the provided inputs and record the tensor operations traced_foo = torch.jit.trace(foo, (torch.rand(3), torch.rand(3))) # `traced_foo` can now be run with the TorchScript interpreter or saved # and loaded in a Python-free environmentПример (отслеживание существующего модуля):
import torch import torch.nn as nn class Net(nn.Module): def __init__(self): super(Net, self).__init__() self.conv = nn.Conv2d(1, 1, 3) def forward(self, x): return self.conv(x) n = Net() example_weight = torch.rand(1, 1, 3, 3) example_forward_input = torch.rand(1, 1, 3, 3) # Trace a specific method and construct `ScriptModule` with # a single `forward` method module = torch.jit.trace(n.forward, example_forward_input) # Trace a module (implicitly traces `forward`) and construct a # `ScriptModule` with a single `forward` method module = torch.jit.trace(n, example_forward_input) - Результат записи автономной функции создаёт
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.jit.trace.html