Rust中如何并行获取ndarray元素可变引用实现并行矩阵乘法
可行的并行矩阵乘法实现方案
以下是基于ndarray和rayon的安全并行矩阵乘法实现,直接并行计算每个元素:
首先确保Cargo.toml中包含正确依赖:
[dependencies] ndarray = { version = "0.15", features = ["rayon"] } rayon = "1.7"
然后实现乘法函数:
use ndarray::{Array2, Axis}; use rayon::iter::ParallelIterator; fn mul(lhs: &Array2<f32>, rhs: &Array2<f32>) -> Array2<f32> { // 检查矩阵维度匹配:左矩阵列数必须等于右矩阵行数 assert_eq!(lhs.dim().1, rhs.dim().0, "Matrix dimensions mismatch for multiplication"); let (n_rows, n_cols) = (lhs.dim().0, rhs.dim().1); let mut result = Array2::zeros((n_rows, n_cols)); // 用ndarray并行迭代器遍历每个元素的索引和可变引用 result.indexed_iter_mut().par_foreach(|((i, j), val)| { // 计算左矩阵第i行与右矩阵第j列的点积 let dot_product = lhs.row(i) .iter() .zip(rhs.column(j).iter()) .map(|(a, b)| a * b) .sum::<f32>(); *val = dot_product; }); result }
关键说明
- 维度校验:必须先验证矩阵乘法的维度合法性,避免无意义的计算触发panic。
- 安全并行:启用
ndarray的rayon特性后,indexed_iter_mut()搭配par_foreach可安全实现多线程并行,每个线程仅负责写入唯一元素位置,无数据竞争风险。 - 点积计算:通过
row(i)和column(j)获取对应行列,用zip配对元素后相乘求和,得到当前位置的乘积值。
如果偏好手动生成索引对的实现方式,也可以用以下安全版本:
use ndarray::Array2; use rayon::iter::{IntoParallelIterator, ParallelIterator}; fn mul(lhs: &Array2<f32>, rhs: &Array2<f32>) -> Array2<f32> { assert_eq!(lhs.dim().1, rhs.dim().0, "Matrix dimensions mismatch for multiplication"); let (n_rows, n_cols) = (lhs.dim().0, rhs.dim().1); let mut result = Array2::zeros((n_rows, n_cols)); // 将结果数组转为可变切片,通过线性索引写入 let result_slice = result.as_mut_slice().unwrap(); // 并行遍历所有(i,j)索引对 (0..n_rows).into_par_iter() .flat_map(|i| (0..n_cols).map(move |j| (i, j))) .for_each(|(i, j)| { let dot_product = lhs.row(i) .iter() .zip(rhs.column(j).iter()) .map(|(a, b)| a * b) .sum::<f32>(); let linear_idx = i * n_cols + j; result_slice[linear_idx] = dot_product; }); result }
内容的提问来源于stack exchange,提问作者stomfaig
相关产品推荐
相关产品推荐

