Spec-Zone.ru › PyTorch 2.14

torch.func.vmap

torch.func.vmap(func, in_dims=0, out_dims=0, randomness='error', *, chunk_size=None) [исходный код]

vmap — это векторизующая функция отображения; vmap(func) возвращает новую функцию, которая применяет отображение func по одному из измерений входных данных. Семантически vmap встраивает отображение в операции PyTorch, вызываемые функцией func, фактически векторизуя эти операции.

vmap полезен для работы с пакетными измерениями: можно написать функцию func, которая выполняется для отдельных примеров, а затем преобразовать её в функцию, способную принимать пакеты примеров с помощью vmap(func). vmap также можно использовать для вычисления пакетных градиентов в сочетании с autograd.

Примечание

torch.vmap() является псевдонимом torch.func.vmap() для удобства. Используйте любой из них на свой выбор.

Параметры:
  • func (функция) – Функция Python, принимающая один или несколько аргументов. Должна возвращать один или несколько тензоров.
  • in_dims (int или вложенная структура) – Задаёт, по какому измерению входных данных следует выполнять отображение. in_dims должна иметь структуру, соответствующую входным данным. Если значение in_dim для конкретного входного аргумента равно None, это означает, что для него нет измерения отображения. По умолчанию: 0.
  • out_dims (int или Tuple[int]) – Задаёт, где в выходных данных должно находиться отображаемое измерение. Если out_dims — это Tuple, он должен содержать по одному элементу для каждого выходного значения. По умолчанию: 0.
  • randomness (str) – Задаёт, должны ли случайные значения в этом vmap быть одинаковыми или различаться в разных пакетах. Если указано ‘different’, случайные значения будут различаться для каждого пакета. Если указано ‘same’, случайные значения будут одинаковыми во всех пакетах. Если указано ‘error’, любые вызовы функций генерации случайных значений приведут к ошибке. По умолчанию: ‘error’. ПРЕДУПРЕЖДЕНИЕ: этот флаг применяется только к операциям PyTorch, использующим случайные значения, и не распространяется на модуль random в Python или генерацию случайных значений в numpy.
  • chunk_size (None или int) – Если None (значение по умолчанию), к входным данным применяется один vmap. Если значение не None, vmap вычисляет chunk_size примеров за один раз. Обратите внимание, что chunk_size=1 эквивалентно вычислению vmap с помощью цикла for. Если при вычислении vmap возникают проблемы с памятью, попробуйте задать chunk_size, отличный от None.
Возвращает:

Возвращает новую «пакетную» функцию. Она принимает те же входные данные, что и func, за исключением того, что у каждого входного значения появляется дополнительное измерение с индексом, заданным в in_dims. Она возвращает те же выходные данные, что и func, за исключением того, что у каждого выходного значения появляется дополнительное измерение с индексом, заданным в out_dims.

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

Callable[[~_P], _R]

Например, vmap() можно использовать для вычисления скалярных произведений пакетами. В PyTorch нет API для пакетного вычисления torch.dot; вместо того чтобы безуспешно искать решение в документации, используйте vmap(), чтобы создать новую функцию.

>>> torch.dot  # [D], [D] -> []
>>> batched_dot = torch.func.vmap(torch.dot)  # [N, D], [N, D] -> [N]
>>> x, y = torch.randn(2, 5), torch.randn(2, 5)
>>> batched_dot(x, y)

vmap() помогает скрыть пакетные измерения, упрощая создание моделей.

>>> batch_size, feature_size = 3, 5
>>> weights = torch.randn(feature_size, requires_grad=True)
>>>
>>> def model(feature_vec):
>>> # Very simple linear model with activation
>>>     return feature_vec.dot(weights).relu()
>>>
>>> examples = torch.randn(batch_size, feature_size)
>>> result = torch.vmap(model)(examples)

vmap() также позволяет векторизовать вычисления, которые ранее было трудно или невозможно выполнять пакетами. Один из примеров — вычисление градиентов высших порядков. Механизм autograd в PyTorch вычисляет vjp (произведения вектора на матрицу Якоби). Для вычисления полной матрицы Якоби некоторой функции f: R^N -> R^N обычно требуется N вызовов autograd.grad, по одному на каждую строку матрицы Якоби. Используя vmap(), можно векторизовать всё вычисление и вычислить матрицу Якоби за один вызов autograd.grad.

>>> # Setup
>>> N = 5
>>> f = lambda x: x**2
>>> x = torch.randn(N, requires_grad=True)
>>> y = f(x)
>>> I_N = torch.eye(N)
>>>
>>> # Sequential approach
>>> jacobian_rows = [torch.autograd.grad(y, x, v, retain_graph=True)[0]
>>>                  for v in I_N.unbind()]
>>> jacobian = torch.stack(jacobian_rows)
>>>
>>> # vectorized gradient computation
>>> def get_vjp(v):
>>>     return torch.autograd.grad(y, x, v)
>>> jacobian = torch.vmap(get_vjp)(I_N)

vmap() также можно вкладывать друг в друга, получая выходные данные с несколькими пакетными измерениями

>>> torch.dot  # [D], [D] -> []
>>> batched_dot = torch.vmap(
...     torch.vmap(torch.dot)
... )  # [N1, N0, D], [N1, N0, D] -> [N1, N0]
>>> x, y = torch.randn(2, 3, 5), torch.randn(2, 3, 5)
>>> batched_dot(x, y)  # tensor of size [2, 3]

Если входные данные сгруппированы не по первому измерению, in_dims задаёт измерение, по которому сгруппированы входные данные

>>> torch.dot  # [N], [N] -> []
>>> batched_dot = torch.vmap(torch.dot, in_dims=1)  # [N, D], [N, D] -> [D]
>>> x, y = torch.randn(2, 5), torch.randn(2, 5)
>>> batched_dot(
...     x, y
... )  # output is [5] instead of [2] if batched along the 0th dimension

Если имеется несколько входных данных, каждый из которых сгруппирован по разным измерениям, in_dims должен быть кортежем, содержащим пакетное измерение для каждого входного значения

>>> torch.dot  # [D], [D] -> []
>>> batched_dot = torch.vmap(torch.dot, in_dims=(0, None))  # [N, D], [D] -> [N]
>>> x, y = torch.randn(2, 5), torch.randn(5)
>>> batched_dot(
...     x, y
... )  # second arg doesn't have a batch dim because in_dim[1] was None

Если входное значение является структурой Python, in_dims должен быть кортежем, содержащим структуру, соответствующую форме входного значения:

>>> f = lambda dict: torch.dot(dict["x"], dict["y"])
>>> x, y = torch.randn(2, 5), torch.randn(5)
>>> input = {"x": x, "y": y}
>>> batched_dot = torch.vmap(f, in_dims=({"x": 0, "y": None},))
>>> batched_dot(input)

По умолчанию выходные данные сгруппированы по первому измерению. Однако с помощью out_dims их можно сгруппировать по любому измерению

>>> f = lambda x: x**2
>>> x = torch.randn(2, 5)
>>> batched_pow = torch.vmap(f, out_dims=1)
>>> batched_pow(x)  # [5, 2]

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

>>> x = torch.randn([2, 5])
>>> def fn(x, scale=4.):
>>>   return x * scale
>>>
>>> batched_pow = torch.vmap(fn)
>>> assert torch.allclose(batched_pow(x), x * 4)
>>> batched_pow(x, scale=x)  # scale is not batched, output has shape [2, 2, 5]

Примечание

vmap не обеспечивает универсальную автоматическую пакетную обработку и не поддерживает последовательности переменной длины «из коробки».

© 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.vmap.html

Spec-Zone.ru

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