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>)[source] -
Отследить модуль и вернуть исполняемый
ScriptModule, который будет оптимизирован с помощью компиляции в реальном времени. Когда модуль передаётся вtorch.jit.trace, выполняется и отслеживается только методforward. С помощьюtrace_module, можно указать словарь имён методов и входных данных для отслеживания (см. аргументinputs) ниже.См.
torch.jit.traceдля получения дополнительной информации об отслеживании.- Параметры:
-
-
mod (torch.nn.Module) – Модуль
torch.nn.Moduleсодержащий методы, имена которых указаны вinputs. Указанные методы будут скомпилированы как часть одногоScriptModule. -
inputs (dict) – Словарь, содержащий примерные входные данные, индексированные по именам методов в
mod. Входные данные будут переданы методам, имена которых соответствуют ключам входных данных во время отслеживания.{ 'forward' : example_forward_input, 'method2': example_method2_input}
-
mod (torch.nn.Module) – Модуль
- Ключевые аргументы:
-
-
check_trace (
bool, optional) – Проверка, что одни и те же входные данные, проходящие через отслеженный код, дают одни и те же выходные данные. По умолчанию:True. Возможно, вы захотите отключить это, если, например, ваша сеть содержит недетерминированные операции или если вы уверены, что сеть правильная, несмотря на ошибку проверки. -
check_inputs (список словарей, необязательно) – Список словарей входных аргументов, которые должны быть использованы для проверки отслеживания по сравнению с ожидаемым результатом. Каждый кортеж эквивалентен набору входных аргументов, которые были бы указаны в
inputs. Для достижения наилучших результатов передайте набор проверочных входных данных, репрезентативных для пространства форм и типов входных данных, которые, как ожидается, увидит сеть. Если не указано, исходныеinputsиспользуются для проверки - check_tolerance (float, необязательно) – Допуск для сравнения чисел с плавающей запятой, используемый в процедуре проверки. Это может использоваться для ослабления строгости проверки в случае, если результаты расходятся численно по известной причине, например, из-за слияния операций.
-
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(Net, self).__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/1.13/generated/torch.jit.trace_module.html