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