如何使用np.matvec实现批量矩阵向量乘法并解决轴参数错误?
使用np.matvec实现矩阵与批量向量的乘法
要实现形状(3,2)的旋转矩阵和形状(2,4,5)的批量向量数组相乘,得到(3,4,5)的结果,关键是通过axes参数为np.matvec的每个输入和输出指定核心轴的对应关系:
核心逻辑
np.matvec的底层广义ufunc签名是(M,N), (N) -> (M),我们需要把批量维度作为广播轴保留,因此要明确:
- 旋转矩阵(
(3,2)):核心轴是自身的两个轴(0,1),对应签名中的(M,N) - 批量向量(
(2,4,5)):核心轴是第0轴(0),对应签名中的(N),后续的(4,5)是批量广播轴 - 输出结果:核心轴是第0轴
(0),对应签名中的(M),批量轴继承自输入的(4,5)
代码实现
import numpy as np # 示例旋转矩阵(形状3x2) rot_matrix = np.array([[np.cos(np.pi/4), -np.sin(np.pi/4)], [np.sin(np.pi/4), np.cos(np.pi/4)], [0, 0]]) # 构造批量向量数组,形状2x4x5 batch_vectors = np.mgrid[:4, :5].astype(np.float64) # 调用np.matvec并指定axes参数 output = np.matvec(rot_matrix, batch_vectors, axes=[(0, 1), (0), (0)]) print(output.shape) # 输出: (3, 4, 5)
错误原因说明
你之前触发的ValueError是因为axes参数必须包含所有输入和输出的核心轴定义(2个输入+1个输出,共3个条目),不能省略输出的核心轴配置,必须明确每个部分的轴映射关系。
内容的提问来源于stack exchange,提问作者DBS4261
相关产品推荐
相关产品推荐

