Spec-Zone.ru › PyTorch 2

Справочник по API torch.func

Преобразования функций

vmap

vmap — это функция векторизации; vmap(func) возвращает новую функцию, которая отображает func по некоторому измерению входных данных.

grad

Оператор grad помогает вычислять градиенты func относительно входных данных, указанных в argnums.

grad_and_value

Возвращает функцию для вычисления кортежа градиента и исходного, или прямого, вычисления.

vjp

Обозначает векторно-якобианское произведение, возвращает кортеж, содержащий результаты func, применённые к primals, и функцию, которая, когда получает cotangents, вычисляет обратный якобиан func относительно primals умноженного на cotangents.

jvp

Обозначает произведение якобиан-вектор, возвращает кортеж, содержащий результат func(*primals) и «якобиан func, вычисленный в primals» умноженный на tangents.

linearize

Возвращает значение func в primals и линейное приближение в primals.

jacrev

Вычисляет якобиан func по аргументу(ам) с индексом argnum с помощью обратного режима автоматической дифференциации.

jacfwd

Вычисляет якобиан func по аргументу(ам) с индексом argnum с помощью прямого режима автоматической дифференциации.

hessian

Вычисляет гессиан func по аргументу(ам) с индексом argnum с помощью стратегии «прямой-через-обратный».

functionalize

functionalize — это преобразование, которое можно использовать для удаления (промежуточных) мутаций и алиасов из функции, сохраняя при этом семантику функции.

Инструменты для работы с модулями torch.nn

В целом, вы можете преобразовать функцию, которая вызывает torch.nn.Module. Например, следующее — пример вычисления якобиана функции, принимающей три значения и возвращающей три значения:

model = torch.nn.Linear(3, 3)

def f(x):
    return model(x)

x = torch.randn(3)
jacobian = jacrev(f)(x)
assert jacobian.shape == (3, 3)

Однако, если вы хотите вычислить якобиан по параметрам модели, необходимо создать функцию, где параметры являются входными данными. Для этого предназначена функция functional_call(): она принимает модуль nn, преобразованную parameters, и входные данные для прохода forward модуля. Она возвращает результат выполнения forward-прохода модуля с заменёнными параметрами.

Вот как мы вычислим якобиан по параметрам

model = torch.nn.Linear(3, 3)

def f(params, x):
    return torch.func.functional_call(model, params, x)

x = torch.randn(3)
jacobian = jacrev(f)(dict(model.named_parameters()), x)
functional_call

Выполняет вызов функции над модулем, заменяя параметры и буферы модуля предоставленными.

stack_module_state

Подготавливает список модулей torch.nn для ансамблирования с помощью vmap().

replace_all_batch_norm_modules_

Вместо этого обновляет root, устанавливая running_mean и running_var в None и устанавливая track_running_stats в False для любого модуля nn.BatchNorm в root.

Если вы ищете информацию о исправлении модулей Batch Norm, следуйте инструкциям здесь

  • Исправление модулей Batch Norm

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/func.api.html

Spec-Zone.ru

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