PyTorch torch.inner Function


Pytorch torch 参考手册PyTorch torch Reference Manual

torch.innerIt is a function in PyTorch used to calculate the inner product of two tensors. For vectors, it is equivalent to the dot product; for higher-dimensional tensors, it computes the inner product along specified dimensions.

Function Definition

torch.inner(input, other, out=None)

Parameters:

  • input(Tensor): The first input tensor.
  • other(Tensor): The second input tensor.
  • out(Tensor, optional): The output tensor.

Return Value:

  • torch.Tensor: Returns the inner product result.

Usage Examples

Example - Vector Inner Product

import torch

# Create two vectors
a = torch.tensor([1.0, 2.0, 3.0])
b = torch.tensor([4.0, 5.0, 6.0])

# Compute the inner product
result = torch.inner(a, b)

print("Vector a:", a)
print("Vector b:", b)
print("Inner product result:", result)

The output result is:

向量 a: tensor([1., 2., 3.])
向量 b: tensor([4., 5., 6.])
内积结果: tensor(32.)

Example - Matrix Inner Product

import torch

# Create two matrices
A = torch.randn(2, 3)
B = torch.randn(2, 3)

# Compute the inner product along the last dimension
result = torch.inner(A, B)
print("A shape:", A.shape)
print("B shape:", B.shape)
print("Result shape:", result.shape)

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions