PyTorch torch.chain_matmul Function
Pytorch torch Reference Manual
torch.chain_matmulis a function in PyTorch used to compute the chain multiplication of multiple matrices. It minimizes computational cost by selecting the optimal matrix multiplication order.
Function Definition
torch.chain_matmul(*matrices, out=None)
Parameters:
matrices(Tensor): The input matrix sequence.out(Tensor, optional): The output tensor.
Return Value:
torch.Tensor: Returns the result of multiplying all matrices.
Usage Example
Example
import torch
# Create multiple matrices
A = torch.randn(10, 20)
B = torch.randn(20, 30)
C = torch.randn(30, 40)
D = torch.randn(40, 50)
# Chained matrix multiplication
result = torch.chain_matmul(A, B, C, D)
print("Matrix A shape:", A.shape)
print("Matrix B shape:", B.shape)
print("Matrix C shape:", C.shape)
print("Matrix D shape:", D.shape)
print("Result shape:", result.shape)
# Create multiple matrices
A = torch.randn(10, 20)
B = torch.randn(20, 30)
C = torch.randn(30, 40)
D = torch.randn(40, 50)
# Chained matrix multiplication
result = torch.chain_matmul(A, B, C, D)
print("Matrix A shape:", A.shape)
print("Matrix B shape:", B.shape)
print("Matrix C shape:", C.shape)
print("Matrix D shape:", D.shape)
print("Result shape:", result.shape)
The output result is:
矩阵 A 形状: torch.Size([10, 20]) 矩阵 B 形状: torch.Size([20, 30]) 矩阵 C 形状: torch.Size([30, 40]) 矩阵 D 形状: torch.Size([40, 50]) 结果形状: torch.Size([10, 50])
Example - Using a List
import torch
# Matrix list
matrices = [torch.randn(10, 20),
torch.randn(20, 30),
torch.randn(30, 40)]
# Use a list as an argument
result = torch.chain_matmul(*matrices)
print("Result shape:", result.shape)
# Matrix list
matrices = [torch.randn(10, 20),
torch.randn(20, 30),
torch.randn(30, 40)]
# Use a list as an argument
result = torch.chain_matmul(*matrices)
print("Result shape:", result.shape)
Other Extensions