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