Spec-Zone.ru › PyTorch 2.14

torch.vmap

torch.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.
Возвращает:

Возвращает новую «пакетную» функцию. Она принимает те же входные данные, что и 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 на пакеты, но будет принимать 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.vmap.html

Spec-Zone.ru

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