torch.jit.trace
-
torch.jit.trace(func, example_inputs=None, 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>, example_kwarg_inputs=None, _store_inputs=True)[source] -
Отследить функцию и вернуть исполняемый или
ScriptFunctionобъект, который будет оптимизирован с помощью компиляции в реальном времени. Отслеживание идеально подходит для кода, который работает только сTensorи списками, словарями и кортежамиTensor.Используя
torch.jit.traceиtorch.jit.trace_module, можно преобразовать существующий модуль или Python-функцию в TorchScriptScriptFunctionилиScriptModule. Вы должны указать примеры входных данных, и мы запустим функцию, записывая операции, выполняемые над всеми тензорами.- Результат записи автономной функции производит
ScriptFunction. - Результат записи
nn.Module.forwardилиnn.ModuleпроизводитScriptModule.
Этот модуль также содержит все параметры, которые были у исходного модуля.
Предупреждение
Отслеживание правильно записывает только функции и модули, которые не зависят от данных (например, не имеют условных операторов по данным в тензорах) и не имеют каких-либо незаписанных внешних зависимостей (например, не выполняют ввод/вывод или не обращаются к глобальным переменным). Отслеживание записывает только операции, выполняемые при запуске данной функции над данными тензорами. Таким образом, возвращённый
ScriptModuleвсегда будет запускать тот же отслеженный граф над любым входом. Это имеет некоторые важные последствия, когда ожидается, что ваш модуль будет выполнять разные наборы операций в зависимости от входных данных и/или состояния модуля. Например,- Отслеживание не будет записывать никакой управляющей логики, такой как операторы if или циклы. Когда это управляющее ветвление является постоянным в вашем модуле, это нормально, и оно часто встраивает решения управляющей логики. Но иногда управляющая логика на самом деле является частью самой модели. Например, рекуррентная сеть представляет собой цикл по (возможно, динамической) длине входной последовательности.
- В возвращённом
ScriptModuleоперации, которые имеют разное поведение вtrainingиevalрежимах, всегда будут вести себя так, как будто они находились в режиме, в котором они находились во время отслеживания, независимо от того, в каком режиме находитсяScriptModule.
В таких случаях отслеживание не подходит, и
scriptingявляется лучшим выбором. Если вы отслеживаете такие модели, вы можете получить неверные результаты при последующих вызовах модели. Отслеживатель будет пытаться выводить предупреждения при выполнении действий, которые могут привести к созданию неправильного отслеживания.- Parameters
-
func (вызываемая функция или torch.nn.Module) – Python-функция или
torch.nn.Module, которая будет запущена сexample_inputs. Аргументыfuncи возвращаемые значения должны быть тензорами или (возможно, вложенными) кортежами, содержащими тензоры. При передаче модуляtorch.jit.trace, запускается и отслеживается только методforward(см.torch.jit.traceдля получения подробностей). - Keyword Arguments
-
-
example_inputs (tuple или torch.Tensor или None, необязательно) – Кортеж примеров входных данных, которые будут переданы функции во время отслеживания. По умолчанию:
None. Должен быть указан либо этот аргумент, либоexample_kwarg_inputs. Полученное отслеживание может быть запущено с входными данными разных типов и форм, предполагая, что отслеженные операции поддерживают эти типы и формы.example_inputsтакже может быть единственным тензором, в этом случае он автоматически оборачивается в кортеж. Если значение равно None, должен быть указанexample_kwarg_inputs. -
check_trace (
bool, необязательно) – Проверка того, что одни и те же входные данные, проходящие через отслеженный код, дают одинаковые выходные данные. По умолчанию:True. Возможно, вам захочется отключить это, если, например, ваша сеть содержит недетерминированные операции или если вы уверены, что сеть верна несмотря на сбой проверки. -
check_inputs (list кортежей, необязательно) – Список кортежей аргументов входных данных, которые следует использовать для проверки отслеживания по сравнению с ожидаемым результатом. Каждый кортеж эквивалентен набору аргументов входных данных, которые были бы указаны в
example_inputs. Для наилучших результатов передайте набор проверочных входных данных, представляющих пространство форм и типов входных данных, которые, как вы ожидаете, будет видеть сеть. Если не указано, используются исходныеexample_inputsдля проверки. - check_tolerance (float, необязательно) – Допуск сравнения с плавающей точкой, который следует использовать в процедуре проверки. Это можно использовать для ослабления строгости проверки в том случае, если результаты расходятся численно по известной причине, например, по слиянию операторов.
-
strict (
bool, необязательно) – запустить отслеживатель в строгом режиме или нет (по умолчанию:True). Отключайте это только тогда, когда вы хотите, чтобы отслеживатель записывал ваши изменяемые контейнерные типы (в настоящее времяlist/dict) и уверены, что используемый вами контейнер — это структураconstantи не используется в качестве условий управляющей логики (if, for). -
example_kwarg_inputs (dict, необязательно) – Этот параметр представляет собой набор ключевых аргументов примеров входных данных, которые будут переданы функции во время отслеживания. По умолчанию:
None. Должен быть указан либо этот аргумент, либоexample_inputs. Словарь будет распакован по именам аргументов отслеживаемой функции. Если ключи словаря не совпадают с именами аргументов отслеживаемой функции, будет вызвано исключение во время выполнения.
-
example_inputs (tuple или torch.Tensor или None, необязательно) – Кортеж примеров входных данных, которые будут переданы функции во время отслеживания. По умолчанию:
- 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().__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/2.1/generated/torch.jit.trace.html