Spec-Zone.ru › PyTorch 2

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

Spec-Zone.ru

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