PyTorch torch.vmap Function
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 vectorizedin_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)
# 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]])
Other Extensions