Spec-Zone.ru › PyTorch 2

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}
Ключевые аргументы
  • check_trace (bool, optional) – Проверка, дают ли одни и те же входные данные при прохождении через прослеженный код одинаковые выходные данные. Значение по умолчанию: True. Возможно, вам захочется отключить эту функцию, если, например, ваша сеть содержит недетерминированные операции или если вы уверены, что сеть верна, несмотря на ошибку проверки.
  • check_inputs (list of dicts, optional) – Список словарей входных аргументов, которые должны использоваться для проверки прослеживания по сравнению с ожидаемым результатом. Каждый кортеж эквивалентен набору входных аргументов, которые были бы указаны в inputs. Для наилучших результатов передайте набор проверочных входных данных, репрезентативных для пространства форм и типов входных данных, которые, как ожидается, увидит сеть. Если не указано, для проверки используются исходные inputs
  • check_tolerance (float, optional) – Допуск сравнения с плавающей точкой, который будет использоваться в процедуре проверки. Это может использоваться для ослабления строгости проверки в том случае, если результаты расходятся численно по известной причине, например, при слиянии операторов.
  • example_inputs_is_kwarg (bool, optional) – Этот параметр указывает, является ли пример входных данных набором ключевых аргументов. Значение по умолчанию: False.
Возвращает

Объект 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

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API