Spec-Zone.ru › PyTorch 2

torch.func.functional_call

torch.func.functional_call(module, parameter_and_buffer_dicts, args, kwargs=None, *, tie_weights=True, strict=False)

Выполняет функциональное обращение к модулю, заменяя параметры и буферы модуля предоставленными.

Примечание

Если модуль имеет активные параметризации, передача значения в аргументе parameter_and_buffer_dicts с именем, установленным в стандартное имя параметра, полностью отключит параметризацию. Если вы хотите применить функцию параметризации к переданному значению, установите ключ как {submodule_name}.parametrizations.{parameter_name}.original.

Примечание

Если модуль выполняет операции на месте с параметрами/буферами, эти операции будут отражены в входных parameter_and_buffer_dicts.

Пример:

>>> a = {'foo': torch.zeros(())}
>>> mod = Foo()  # does self.foo = self.foo + 1
>>> print(mod.foo)  # tensor(0.)
>>> functional_call(mod, a, torch.ones(()))
>>> print(mod.foo)  # tensor(0.)
>>> print(a['foo'])  # tensor(1.)

Примечание

Если модуль имеет связанные веса, то соблюдение привязки в functional_call определяется флагом tie_weights.

Пример:

>>> a = {'foo': torch.zeros(())}
>>> mod = Foo()  # has both self.foo and self.foo_tied which are tied. Returns x + self.foo + self.foo_tied
>>> print(mod.foo)  # tensor(1.)
>>> mod(torch.zeros(()))  # tensor(2.)
>>> functional_call(mod, a, torch.zeros(()))  # tensor(0.) since it will change self.foo_tied too
>>> functional_call(mod, a, torch.zeros(()), tie_weights=False)  # tensor(1.)--self.foo_tied is not updated
>>> new_a = {'foo', torch.zeros(()), 'foo_tied': torch.zeros(())}
>>> functional_call(mod, new_a, torch.zeros()) # tensor(0.)

Пример передачи нескольких словарей

a = ({'weight': torch.ones(1, 1)}, {'buffer': torch.zeros(1)})  # two separate dictionaries
mod = nn.Bar(1, 1)  # return self.weight @ x + self.buffer
print(mod.weight)  # tensor(...)
print(mod.buffer)  # tensor(...)
x = torch.randn((1, 1))
print(x)
functional_call(mod, a, x)  # same as x
print(mod.weight)  # same as before functional_call

И вот пример применения преобразования grad над параметрами модели.

import torch
import torch.nn as nn
from torch.func import functional_call, grad

x = torch.randn(4, 3)
t = torch.randn(4, 3)
model = nn.Linear(3, 3)

def compute_loss(params, x, t):
    y = functional_call(model, params, x)
    return nn.functional.mse_loss(y, t)

grad_weights = grad(compute_loss)(dict(model.named_parameters()), x, t)

Примечание

Если пользователю не требуется отслеживание grad за пределами преобразований grad, он может отсоединить все параметры для лучшей производительности и использования памяти.

Пример:

>>> detached_params = {k: v.detach() for k, v in model.named_parameters()}
>>> grad_weights = grad(compute_loss)(detached_params, x, t)
>>> grad_weights.grad_fn  # None--it's not tracking gradients outside of grad

Это означает, что пользователь не может вызвать grad_weight.backward(). Однако, если отслеживание autograd не требуется за пределами преобразований, это приведет к меньшему использованию памяти и более высокой скорости.

Параметры
  • module (torch.nn.Module) – модуль, который нужно вызвать
  • parameters_and_buffer_dicts (Dict[str, Tensor] or tuple of Dict[str, Tensor]) – параметры, которые будут использоваться при вызове модуля. Если передается кортеж словарей, они должны иметь разные ключи, чтобы все словари можно было использовать вместе
  • args (Any or tuple) – аргументы, которые будут переданы при вызове модуля. Если не кортеж, считается единственным аргументом.
  • kwargs (dict) – ключевые аргументы, которые будут переданы при вызове модуля
  • tie_weights (bool, optional) – Если True, параметры и буферы, привязанные в исходной модели, будут обрабатываться как привязанные в перепараметризованной версии. Поэтому, если True и переданы разные значения для привязанных параметров и буферов, произойдет ошибка. Если False, исходные привязанные параметры и буферы не будут учитываться, если значения, переданные для обоих весов, не совпадают. По умолчанию: True.
  • strict (bool, optional) – Если True, параметры и буферы, переданные, должны совпадать с параметрами и буферами в исходном модуле. Таким образом, если True и есть какие-либо отсутствующие или неожиданные ключи, произойдет ошибка. По умолчанию: False.
Возвращает

результат вызова module.

Тип возвращаемого значения

Any

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.func.functional_call.html

Spec-Zone.ru

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