PyTorch torch.chain_matmul Function


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

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)

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions