PyTorch torch.vmap Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.vmapIt is a function in PyTorch used for vector mapping. It takes a function as input and returns a new function that can automatically operate on batched tensors, similar to vmap in JAX.

Function Definition

torch.vmap(func, in_dims, out_dims, randomness)

Parameter Description

  • func: the function to be vectorized
  • in_dims: the batch dimension of the input tensor (optional)
  • out_dims: the batch dimension of the output tensor (optional)
  • randomness: random behavior, optional "error", "different", "same"

Usage Example

Example

import torch

# Define a simple function
def simple_func(x):
    return x * 2 + 1

# Use vmap to vectorize the function
vectorized_func = torch.vmap(simple_func)

# Batched input (batch dimension is 0)
batch_input = torch.randn(4, 3)

# Apply the vectorized function
output = vectorized_func(batch_input)

print("Input shape:", batch_input.shape)
print("Output shape:", output.shape)
print("Output:")
print(output)

The output result is:

输入形状: torch.Size([4, 3])
输出形状: torch.Size([4, 3])
输出:
tensor([[ 0.2345,  1.5678, -0.3456],
        [ 2.1234, -1.2345,  0.5678],
        [-0.8765,  1.2345,  2.3456],
        [ 1.5678,  0.1234, -1.2345]])

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions