Rust中如何用Range替代索引数组实现矩阵乘向量?
改用Range索引实现矩阵与向量的栈上乘积
你可以利用Rust标准库中的core::array::from_fn(Rust 1.55及以上版本可用)结合0..3这样的range来简化代码,直接生成固定长度的结果数组,完全替代手动定义的索引数组。
修改后的代码如下:
fn main() { let matrix: [[f64; 3]; 3] = [ [1.0, 0.0, 0.0], [-1.0, 1.0, 0.0], [1.0, 0.0, -1.0], ]; let vector: [f64; 3] = [3.0, -2.0, 1.0]; // 用from_fn生成结果数组,i为矩阵的行索引 let result: [f64; 3] = core::array::from_fn(|i| { // 通过0..3遍历列索引j,计算乘积后求和 (0..3).map(|j| matrix[i][j] * vector[j]).sum() }); println!("{:?}", result); }
代码说明
core::array::from_fn是专门用于生成固定长度数组的工具,它接收一个闭包,闭包参数是数组的索引(这里对应矩阵的行索引i),返回对应位置的元素值。- 计算每行的和时,用
0..3这个range遍历列索引j,通过map迭代器计算矩阵元素与向量对应元素的乘积,最后调用sum()完成累加,逻辑和原代码完全一致,但去掉了冗余的索引数组。
如果你偏好使用矩阵行的迭代器风格,也可以用下面的写法(需要将迭代结果转换为固定数组):
fn main() { let matrix: [[f64; 3]; 3] = [ [1.0, 0.0, 0.0], [-1.0, 1.0, 0.0], [1.0, 0.0, -1.0], ]; let vector: [f64; 3] = [3.0, -2.0, 1.0]; let result: [f64; 3] = matrix .iter() .map(|row| { row.iter() .enumerate() .map(|(j, &val)| val * vector[j]) .sum() }) .collect::<Vec<_>>() .try_into() .unwrap(); println!("{:?}", result); }
不过第一种from_fn的写法更高效,因为它直接在栈上构造目标数组,不需要中间的Vec临时存储。
内容的提问来源于stack exchange,提问作者Bernhard Bodenstorfer
相关产品推荐
相关产品推荐

