如何用Rust的nalgebra实现类似numpy.dot的矩阵乘法?
问题描述
最近开始使用nalgebra,但始终无法实现常规矩阵乘法,想要的效果类似Python中numpy.dot的操作:
import numpy as np arr1 = np.array([[1, 2, 3], [4, 5, 6]]) arr2 = np.array([[7, 8], [9, 10], [11, 12]]) print(arr1.dot(arr2)) # [[ 58 64][139 154]]
但在Rust中用nalgebra尝试时得到错误结果:
let matrix1 = Matrix3x2::from_vec(vec![1, 2, 3, 4, 5, 6]); let matrix2 = Matrix2x3::from_vec(vec![7, 8, 9, 10, 11, 12]); println!("{:?}", matrix1); // [[1, 2, 3], [4, 5, 6]] println!("{:?}", matrix2); // [[7, 8], [9, 10], [11, 12]] println!("{:?}", matrix1 * matrix2); // [[39, 54, 69], [49, 68, 87], [59, 82, 105]]
尝试了mul、cross、dot等函数,要么无法运行要么结果错误。请问忽略了什么?
问题原因与解决方法
核心问题是矩阵维度声明错误:
nalgebra中MatrixMxN的命名规则是Matrix<行数>x<列数>,比如Matrix2x3代表2行3列的矩阵,Matrix3x2代表3行2列的矩阵。- 你需要的第一个矩阵是2行3列(对应numpy的
arr1),但错误地用了Matrix3x2(3行2列);第二个矩阵是3行2列(对应numpy的arr2),却用了Matrix2x3(2行3列)。维度不匹配导致乘法结果完全错误。
正确的代码写法:
use nalgebra::{Matrix2x3, Matrix3x2}; fn main() { // 2行3列的矩阵,对应numpy的arr1 let matrix1 = Matrix2x3::from_vec(vec![1, 2, 3, 4, 5, 6]); // 3行2列的矩阵,对应numpy的arr2 let matrix2 = Matrix3x2::from_vec(vec![7, 8, 9, 10, 11, 12]); println!("{:?}", matrix1); // [[1, 2, 3], [4, 5, 6]] println!("{:?}", matrix2); // [[7, 8], [9, 10], [11, 12]] println!("{:?}", matrix1 * matrix2); // [[58, 64], [139, 154]] }
额外说明:
nalgebra中*运算符就是常规矩阵乘法,只要维度匹配(第一个矩阵的列数等于第二个矩阵的行数)就能得到正确结果。dot方法是针对向量的点积,不是矩阵乘法;cross是叉乘,仅适用于3维向量;mul和*作用一致,都是矩阵乘法。
内容的提问来源于stack exchange,提问作者Ayush Garg
相关产品推荐
相关产品推荐

