PyTorch torch.cross Function


Pytorch torch 参考手册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)

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions