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