Spec-Zone.ru › PyTorch 2

Миграция с 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

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, пожалуйста, попробуйте использовать эквиваленты 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()

Утилиты модулей нейронных сетей

Мы изменили 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

Spec-Zone.ru

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