torch.jit.trace_module
-
torch.jit.trace_module(mod, 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>, example_inputs_is_kwarg=False, _store_inputs=True)[source] -
Проследите модуль и верните исполняемый
ScriptModule, который будет оптимизирован с помощью компиляции just-in-time. Когда модуль передается вtorch.jit.trace, выполняется и прослеживается только методforward. С помощьюtrace_module, вы можете указать словарь имен методов и входных данных для прослеживания (см. аргументinputs) ниже.См.
torch.jit.traceдля получения дополнительной информации о прослеживании.- Параметры
-
-
mod (torch.nn.Module) – A
torch.nn.Moduleсодержащий методы, имена которых указаны вinputs. Указанные методы будут скомпилированы как часть одногоScriptModule. -
inputs (dict) – Словарь, содержащий образцы входных данных, индексированные по именам методов в
mod. Входные данные будут переданы методам, имена которых соответствуют ключам входных данных во время прослеживания.{ 'forward' : example_forward_input, 'method2': example_method2_input}
-
mod (torch.nn.Module) – A
- Ключевые аргументы
-
-
check_trace (
bool, optional) – Проверка, дают ли одни и те же входные данные при прохождении через прослеженный код одинаковые выходные данные. Значение по умолчанию:True. Возможно, вам захочется отключить эту функцию, если, например, ваша сеть содержит недетерминированные операции или если вы уверены, что сеть верна, несмотря на ошибку проверки. -
check_inputs (list of dicts, optional) – Список словарей входных аргументов, которые должны использоваться для проверки прослеживания по сравнению с ожидаемым результатом. Каждый кортеж эквивалентен набору входных аргументов, которые были бы указаны в
inputs. Для наилучших результатов передайте набор проверочных входных данных, репрезентативных для пространства форм и типов входных данных, которые, как ожидается, увидит сеть. Если не указано, для проверки используются исходныеinputs - check_tolerance (float, optional) – Допуск сравнения с плавающей точкой, который будет использоваться в процедуре проверки. Это может использоваться для ослабления строгости проверки в том случае, если результаты расходятся численно по известной причине, например, при слиянии операторов.
-
example_inputs_is_kwarg (
bool, optional) – Этот параметр указывает, является ли пример входных данных набором ключевых аргументов. Значение по умолчанию:False.
-
check_trace (
- Возвращает
-
Объект
ScriptModuleс единственным методомforward, содержащим прослеженный код. Когдаfuncявляетсяtorch.nn.Module, возвращаемыйScriptModuleбудет иметь тот же набор подмодулей и параметров, что иfunc.
Пример (прослеживание модуля с несколькими методами):
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) def weighted_kernel_sum(self, weight): return weight * self.conv.weight 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) # Trace specific methods on a module (specified in `inputs`), constructs # a `ScriptModule` with `forward` and `weighted_kernel_sum` methods inputs = {'forward' : example_forward_input, 'weighted_kernel_sum' : example_weight} module = torch.jit.trace_module(n, inputs)
© 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_module.html