如何在矩阵乘法算法中满足Rayon的Trait约束?
矩阵乘法并行化的Rayon错误问题解决
问题描述
我用函数式方法实现了矩阵乘法,代码如下:
struct Matrix { vals: Vec<i32>, rows: usize, cols: usize, } fn mul(a: &Matrix, b: &Matrix) -> Matrix { assert_eq!(a.cols, b.rows); let vals = a .vals .par_chunks(a.cols) .map(|row| { (0..b.cols).map(|i| { row .iter() .zip(b.vals.iter().skip(i).step_by(b.cols)) .map(|(a, b)| a * b) .sum() }) }) .flatten() .collect(); Matrix { vals, rows: a.rows, cols: b.cols, } }
尝试用Rayon 1.8.1的par_chunks替代chunks并行化算法,让每行在单独线程中构建。用chunks时正常,但用par_chunks时出现两个错误:
error[E0277]: the trait bound `std::iter::Map<std::ops::Range<usize>, {closure@src\main.rs:145:23: 145:26}>: rayon::iter::ParallelIterator` is not satisfied --> src\main.rs:153:6 | 153 | .flatten() | ^^^^^^^ the trait `rayon::iter::ParallelIterator` is not implemented for `std::iter::Map<std::ops::Range<usize>, {closure@src\main.rs:145:23: 145:26}>`
error[E0599]: the method `collect` exists for struct `Flatten<Map<Chunks<'_, i32>, {closure@main.rs:144:10}>>`, but its trait bounds were not satisfied --> src\main.rs:154:6 | 141 | let vals = a | ______________- 142 | | .vals 143 | | .par_chunks(a.cols) 144 | | .map(|row| { ... | 153 | | .flatten() 154 | | .collect(); | | -^^^^^^^ method cannot be called due to unsatisfied trait bounds | |_____| | | ::: C:\rust\.rustup\toolchains\stable-x86_64-pc-windows-gnu\lib/rustlib/src/rust\library\core\src\iter\adapters\map.rs:62:1 | 62 | pub struct Map<I, F> { | -------------------- doesn't satisfy `_: IntoParallelIterator` | ::: C:\rust\.cargo\registry\src\index.crates.io-6f17d22bba15001f\rayon-1.8.1\src\iter\flatten.rs:11:1 | 11 | pub struct Flatten<I: ParallelIterator> { | --------------------------------------- | | | doesn't satisfy `_: Iterator` | doesn't satisfy `_: ParallelIterator`
已引入Rayon的prelude,不清楚问题所在,怀疑和多线程共享第二个矩阵有关,但该矩阵是只读且在作用域内,想知道哪里出错以及如何解读错误信息。
错误原因解读
- 第一个错误的核心:Rayon的
flatten()方法要求被扁平化的迭代器必须是ParallelIterator类型,但你在map闭包里返回的是标准库的普通Map迭代器((0..b.cols).map(...)),它并没有实现ParallelIteratortrait,所以Rayon无法对其进行扁平化操作。 - 第二个错误是连锁反应:因为
flatten()的 trait 约束不满足,导致后续的collect()也无法调用——Rayon的并行迭代器和标准库迭代器的collect()依赖的 trait 不同,这里的迭代器链已经因为flatten()的问题既不满足标准库Iterator也不满足RayonParallelIterator的约束。
注意:你怀疑的共享矩阵问题并不存在,b是只读引用,Rayon可以安全地在多线程中共享只读数据。
修复方案
有两种常见的修复方式:
方式一:将内部迭代器转为并行迭代器
把闭包里的普通迭代器换成Rayon的并行迭代器,用into_par_iter()替代普通迭代器构造方式:
fn mul(a: &Matrix, b: &Matrix) -> Matrix { assert_eq!(a.cols, b.rows); let vals = a .vals .par_chunks(a.cols) .map(|row| { // 将Range转为并行迭代器 (0..b.cols).into_par_iter().map(|i| { row .iter() .zip(b.vals.iter().skip(i).step_by(b.cols)) .map(|(a, b)| a * b) .sum() }) }) .flatten() .collect(); Matrix { vals, rows: a.rows, cols: b.cols, } }
这里用into_par_iter()把Range转为并行迭代器,闭包里返回的就是ParallelIterator,满足flatten()的约束。
方式二:在闭包里先收集为Vec,避免嵌套迭代器
如果不需要内部也并行,可以把闭包里的迭代器先收集成Vec<i32>,这样map返回的是具体集合而非迭代器,Rayon的flatten()可以直接处理:
fn mul(a: &Matrix, b: &Matrix) -> Matrix { assert_eq!(a.cols, b.rows); let vals = a .vals .par_chunks(a.cols) .map(|row| { (0..b.cols).map(|i| { row .iter() .zip(b.vals.iter().skip(i).step_by(b.cols)) .map(|(a, b)| a * b) .sum() }) .collect::<Vec<i32>>() // 先收集成Vec,消除嵌套迭代器 }) .flatten() .collect(); Matrix { vals, rows: a.rows, cols: b.cols, } }
这种方式下,map返回的是Vec<i32>,Rayon会自动把它当成可并行迭代的序列,flatten()可以正常工作。
内容的提问来源于stack exchange,提问作者Sun of A beach
相关产品推荐
相关产品推荐

