PyTorch torch.cross Function
PyTorch torch Reference Manual
torch.crossIt is a function in PyTorch used to compute the cross product of two 3-dimensional vectors (or batches of 3-dimensional vectors). The cross product produces a new vector perpendicular to both input vectors.
Function Definition
torch.cross(input, other, dim=-1)
Usage Example
Example
import torch
# Cross product of two 3-dimensional vectors
a = torch.tensor([1, 0, 0])
b = torch.tensor([0, 1, 0])
c = torch.cross(a, b)
print("a:", a)
print("b:", b)
print("a x b:", c)
# Output: tensor([0, 0, 1])
# Compute cross product in batch
a = torch.tensor([[1, 0, 0], [0, 1, 0]])
b = torch.tensor([[0, 1, 0], [1, 0, 0]])
result = torch.cross(a, b)
print("Batch cross product:")
print(result)
# tensor([[0, 0, 1],
# [0, 0, -1]])
# Specify the dimension
a = torch.randn(3, 4, 3)
b = torch.randn(3, 4, 3)
result = torch.cross(a, b, dim=2)
print("Cross product shape with specified dimension:", result.shape)
# Cross product of two 3-dimensional vectors
a = torch.tensor([1, 0, 0])
b = torch.tensor([0, 1, 0])
c = torch.cross(a, b)
print("a:", a)
print("b:", b)
print("a x b:", c)
# Output: tensor([0, 0, 1])
# Compute cross product in batch
a = torch.tensor([[1, 0, 0], [0, 1, 0]])
b = torch.tensor([[0, 1, 0], [1, 0, 0]])
result = torch.cross(a, b)
print("Batch cross product:")
print(result)
# tensor([[0, 0, 1],
# [0, 0, -1]])
# Specify the dimension
a = torch.randn(3, 4, 3)
b = torch.randn(3, 4, 3)
result = torch.cross(a, b, dim=2)
print("Cross product shape with specified dimension:", result.shape)
Other Extensions