Spec-Zone.ru › PyTorch 2.14

torch.func.functional_call

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

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

Примечание

Если в модуле активны параметризации, передача значения в аргумент 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, он может отсоединить все параметры для повышения производительности и сокращения использования памяти.

Пример:

>>> 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

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

Spec-Zone.ru

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