Spec-Zone.ru › PyTorch 2.14

Переход с 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

torch.vmap() или torch.func.vmap()

functorch.grad

torch.func.grad()

functorch.vjp

torch.func.vjp()

functorch.jvp

torch.func.jvp()

functorch.jacrev

torch.func.jacrev()

functorch.jacfwd

torch.func.jacfwd()

functorch.hessian

torch.func.hessian()

functorch.functionalize

torch.func.functionalize()

Кроме того, если вы используете API torch.autograd.functional, попробуйте вместо них соответствующие API torch.func. Преобразования функций torch.func лучше компонуются и во многих случаях работают быстрее.

API torch.autograd.functional

API torch.func (начиная с PyTorch 2.0)

torch.autograd.functional.vjp()

torch.func.grad() или torch.func.vjp()

torch.autograd.functional.jvp()

torch.func.jvp()

torch.autograd.functional.jacobian()

torch.func.jacrev() или torch.func.jacfwd()

torch.autograd.functional.hessian()

torch.func.hessian()

Утилиты для модулей 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

Spec-Zone.ru

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