Spec-Zone.ru › PyTorch 1

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-функцию в TorchScript ScriptFunction или 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 также может быть единственным тензором, в этом случае он автоматически упаковывается в кортеж.
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).
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

Spec-Zone.ru

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