Справочник по API torch.func
Преобразования функций
vmap
| vmap — это функция векторизации; |
grad
| Оператор |
grad_and_value
| Возвращает функцию для вычисления кортежа градиента и исходного, или прямого, вычисления. |
vjp
| Обозначает векторно-якобианское произведение, возвращает кортеж, содержащий результаты |
jvp
| Обозначает произведение якобиан-вектор, возвращает кортеж, содержащий результат |
linearize
| Возвращает значение |
jacrev
| Вычисляет якобиан |
jacfwd
| Вычисляет якобиан |
hessian
| Вычисляет гессиан |
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 для ансамблирования с помощью |
replace_all_batch_norm_modules_
| Вместо этого обновляет |
Если вы ищете информацию о исправлении модулей 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