Переход с functorch на torch.func
Создано: 11 июня 2025 г. | Последнее обновление: 11 июня 2025 г.
torch.func, ранее известный как «functorch», — это похожие на JAX компонуемые преобразования функций для PyTorch.
functorch начинался как внешняя библиотека в репозитории pytorch/functorch. Наша цель всегда заключалась в том, чтобы перенести functorch непосредственно в PyTorch и предоставить его в качестве основной библиотеки PyTorch.
В качестве заключительного этапа переноса мы решили перейти от пакета верхнего уровня (functorch) к включению в PyTorch, чтобы отразить интеграцию преобразований функций непосредственно в ядро PyTorch. Начиная с PyTorch 2.0, мы объявляем import functorch устаревшим и просим пользователей перейти на новейшие API, поддержку которых мы будем продолжать. import functorch останется доступным для обратной совместимости в течение нескольких выпусков.
Преобразования функций
Следующие API являются прямой заменой перечисленных ниже API functorch. Они полностью обратно совместимы.
API functorch | API PyTorch (начиная с PyTorch 2.0) |
|---|---|
functorch.vmap | |
functorch.grad | |
functorch.vjp | |
functorch.jvp | |
functorch.jacrev | |
functorch.jacfwd | |
functorch.hessian | |
functorch.functionalize |
Кроме того, если вы используете API torch.autograd.functional, попробуйте вместо них соответствующие API torch.func. Преобразования функций torch.func лучше компонуются и во многих случаях работают быстрее.
API torch.autograd.functional | API torch.func (начиная с PyTorch 2.0) |
|---|---|
Утилиты для модулей NN
Мы изменили API для применения преобразований функций к модулям NN, чтобы лучше соответствовать философии проектирования PyTorch. Новый API отличается, поэтому внимательно прочитайте этот раздел.
functorch.make_functional
torch.func.functional_call() заменяет functorch.make_functional и functorch.make_functional_with_buffers. Однако это не прямая замена.
Если вы спешите, воспользуйтесь вспомогательными функциями из этого gist, которые имитируют поведение functorch.make_functional и functorch.make_functional_with_buffers. Мы рекомендуем напрямую использовать torch.func.functional_call(), поскольку это более явный и гибкий API.
В частности, functorch.make_functional возвращает функциональный модуль и параметры. Функциональный модуль принимает параметры и входные данные модели в качестве аргументов. torch.func.functional_call() позволяет выполнить прямой проход существующего модуля с новыми параметрами, буферами и входными данными.
Вот пример вычисления градиентов параметров модели с помощью functorch и torch.func:
# ---------------
# using functorch
# ---------------
import torch
import functorch
inputs = torch.randn(64, 3)
targets = torch.randn(64, 3)
model = torch.nn.Linear(3, 3)
fmodel, params = functorch.make_functional(model)
def compute_loss(params, inputs, targets):
prediction = fmodel(params, inputs)
return torch.nn.functional.mse_loss(prediction, targets)
grads = functorch.grad(compute_loss)(params, inputs, targets)
# ------------------------------------
# using torch.func (as of PyTorch 2.0)
# ------------------------------------
import torch
inputs = torch.randn(64, 3)
targets = torch.randn(64, 3)
model = torch.nn.Linear(3, 3)
params = dict(model.named_parameters())
def compute_loss(params, inputs, targets):
prediction = torch.func.functional_call(model, params, (inputs,))
return torch.nn.functional.mse_loss(prediction, targets)
grads = torch.func.grad(compute_loss)(params, inputs, targets)
А вот пример вычисления якобианов параметров модели:
# --------------- # using functorch # --------------- import torch import functorch inputs = torch.randn(64, 3) model = torch.nn.Linear(3, 3) fmodel, params = functorch.make_functional(model) jacobians = functorch.jacrev(fmodel)(params, inputs) # ------------------------------------ # using torch.func (as of PyTorch 2.0) # ------------------------------------ import torch from torch.func import jacrev, functional_call inputs = torch.randn(64, 3) model = torch.nn.Linear(3, 3) params = dict(model.named_parameters()) # jacrev computes jacobians of argnums=0 by default. # We set it to 1 to compute jacobians of params jacobians = jacrev(functional_call, argnums=1)(model, params, (inputs,))
Обратите внимание: для экономии памяти важно хранить только одну копию параметров. model.named_parameters() не копирует параметры. Если во время обучения модели вы обновляете её параметры на месте, то nn.Module, то есть ваша модель, содержит единственную копию параметров, и всё работает правильно.
Однако если вы хотите хранить параметры в словаре и обновлять их не на месте, то существуют две копии параметров: одна в словаре, а другая в model. В этом случае следует преобразовать model в устройство meta с помощью model.to('meta'), чтобы оно не занимало память.
functorch.combine_state_for_ensemble
Используйте torch.func.stack_module_state() вместо functorch.combine_state_for_ensemble. torch.func.stack_module_state() возвращает два словаря: один со стеком параметров, а другой — со стеком буферов. Затем их можно использовать с torch.vmap() и torch.func.functional_call() для создания ансамбля.
Например, вот пример создания ансамбля на основе очень простой модели:
import torch
num_models = 5
batch_size = 64
in_features, out_features = 3, 3
models = [torch.nn.Linear(in_features, out_features) for i in range(num_models)]
data = torch.randn(batch_size, 3)
# ---------------
# using functorch
# ---------------
import functorch
fmodel, params, buffers = functorch.combine_state_for_ensemble(models)
output = functorch.vmap(fmodel, (0, 0, None))(params, buffers, data)
assert output.shape == (num_models, batch_size, out_features)
# ------------------------------------
# using torch.func (as of PyTorch 2.0)
# ------------------------------------
import copy
# Construct a version of the model with no memory by putting the Tensors on
# the meta device.
base_model = copy.deepcopy(models[0])
base_model.to('meta')
params, buffers = torch.func.stack_module_state(models)
# It is possible to vmap directly over torch.func.functional_call,
# but wrapping it in a function makes it clearer what is going on.
def call_single_model(params, buffers, data):
return torch.func.functional_call(base_model, (params, buffers), (data,))
output = torch.vmap(call_single_model, (0, 0, None))(params, buffers, data)
assert output.shape == (num_models, batch_size, out_features)
functorch.compile
Мы больше не поддерживаем functorch.compile (также известный как AOTAutograd) в качестве фронтенда для компиляции в PyTorch; мы интегрировали AOTAutograd в механизм компиляции PyTorch. Если вы используете этот API, воспользуйтесь вместо него torch.compile().
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/func.migrating.html