Миграция с functorch на torch.func
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, пожалуйста, попробуйте использовать эквиваленты torch.func вместо этого. Преобразования функций torch.func в большинстве случаев более композируемы и производительны.
API torch.autograd.functional | API torch.func (начиная с PyTorch 2.0) |
|---|---|
Утилиты модулей нейронных сетей
Мы изменили API, чтобы применять преобразования функций к модулям нейронных сетей, чтобы они лучше соответствовали философии проектирования 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 так, чтобы он не удерживал память, преобразуя его в метаустройство с помощью 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. Если вы пользователь, пожалуйста, используйте torch.compile() вместо этого.
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/func.migrating.html