Rust中如何实现返回带Item trait约束的并行迭代器?
矩阵通用接口的并行迭代问题解决方法
问题背景
我正在为不同类型的矩阵编写通用接口,支持可变迭代行并修改。现有两种矩阵类型:
struct NdArrayMatrix { matrix: Array2<f32>, } struct ByteMatrix<'a> { data: &'a mut [u8], rows: usize, cols: usize, }
NdArrayMatrix是内存存储的矩阵,ByteMatrix基于MMap实现内存映射(已省略细节)。
首先定义ReadWrite trait统一读写操作:
trait ReadWrite { fn rw_read(&self, i: usize, j: usize) -> f32; fn rw_write(&mut self, i: usize, j: usize, val: f32); }
接着创建Sliceable trait,返回rayon::iter::IndexedParallelIterator:
trait Sliceable<'a> { type Output: IndexedParallelIterator; fn rows_par_iter(&'a mut self ) -> Self::Output; }
但在泛型上下文使用时遇到问题,比如:
fn<'a, T> slice_and_write(matrix: T) where T: Sliceable<'a> { T.rows_par_iter() .map(|mut row| { row.rw_write(...); }) ... }
此时row未实现ReadWrite,报错。于是尝试定义RwIterator trait:
trait RwIterator: IndexedParallelIterator { type Item: ReadWrite; }
修改Sliceable trait:
trait Sliceable<'a> { type Output: RwIterator; fn rows_par_iter(&'a mut self ) -> Self::Output; }
但仍报错:
| row.rw_write(...); | ^^^^^^^^ method not found in `<<T as Sliceable<'a>>::Output as ParallelIterator>::Item`
怀疑map仅受ParallelIterator约束,无法利用RwIterator的约束。
问题:有没有办法解决这个问题,或者其他替代实现方式?
最小可复现代码
use ndarray::Array2; use rayon::prelude::*; use ndarray::Axis; use ndarray::parallel::Parallel; use ndarray::Dim; use ndarray::iter::AxisIterMut; use rayon::iter::ParallelIterator; use ndarray::ViewRepr; use ndarray::ArrayBase; struct NdArrayMatrix { matrix: Array2<f32>, } impl NdArrayMatrix { pub fn new() -> Self { let matrix = Array2::zeros((10, 10)); Self { matrix, } } } trait ReadWrite { fn rw_read(&self, i: usize, j: usize) -> f32; fn rw_write(&mut self, i: usize, j: usize, val: f32); } impl ReadWrite for NdArrayMatrix { fn rw_read(&self, i: usize, j: usize) -> f32 { self.matrix[[i, j]] } fn rw_write(&mut self, i: usize, j: usize, val: f32) { self.matrix[[i, j]] = val; } } impl ReadWrite for ArrayBase<ViewRepr<&mut f32>, Dim<[usize; 1]>> { fn rw_read(&self, i: usize, j: usize) -> f32 { self[j] } fn rw_write(&mut self, i: usize, j: usize, val: f32) { self[j] = val; } } trait RwIterator: IndexedParallelIterator { type Item: ReadWrite; } impl<'a> RwIterator for Parallel<AxisIterMut<'a, f32, Dim<[usize; 1]>>> { type Item = ArrayBase<ViewRepr<&'a mut f32>, Dim<[usize; 1]>>; } trait Sliceable<'a> { type Output: RwIterator; fn rows_par_iter(&'a mut self ) -> Self::Output; } impl<'a> Sliceable<'a> for NdArrayMatrix { type Output = Parallel<AxisIterMut<'a, f32, Dim<[usize; 1]>>>; fn rows_par_iter(&'a mut self) -> Self::Output { self.matrix .axis_iter_mut(Axis(0)) .into_par_iter() } } fn main() { let mut matrix: NdArrayMatrix = NdArrayMatrix::new(); test(matrix); } fn test<'a, T> (matrix: T) where T: Sliceable<'a> + ReadWrite { matrix.rows_par_iter() .map(|mut row| { row.rw_write(0, 0, 0.0); }).count(); }
解决方法
核心问题是间接通过RwIterator约束迭代器Item的方式,无法被ParallelIterator的方法(如map)识别。我们需要在Sliceable trait中直接明确迭代器的Item必须实现ReadWrite。
步骤1:修改Sliceable trait
直接关联迭代器的Item类型,并约束其实现ReadWrite:
use rayon::iter::IndexedParallelIterator; trait Sliceable<'a> { type IterItem: ReadWrite; type Output: IndexedParallelIterator<Item = Self::IterItem>; fn rows_par_iter(&'a mut self) -> Self::Output; }
步骤2:修正Sliceable实现
为NdArrayMatrix实现Sliceable时,显式指定关联类型:
impl<'a> Sliceable<'a> for NdArrayMatrix { type IterItem = ArrayBase<ViewRepr<&'a mut f32>, Dim<[usize; 1]>>; type Output = Parallel<AxisIterMut<'a, f32, Dim<[usize; 1]>>>; fn rows_par_iter(&'a mut self) -> Self::Output { self.matrix.axis_iter_mut(Axis(0)).into_par_iter() } }
步骤3:修正泛型函数参数
rows_par_iter需要&mut self,因此泛型函数的参数必须是可变引用:
fn test<'a, T>(matrix: &'a mut T) where T: Sliceable<'a> + ReadWrite { matrix.rows_par_iter() .map(|mut row| { row.rw_write(0, 0, 0.0); }).count(); }
步骤4:修正main函数调用
传递可变引用给test函数:
fn main() { let mut matrix: NdArrayMatrix = NdArrayMatrix::new(); test(&mut matrix); }
完整修正代码
use ndarray::Array2; use rayon::prelude::*; use ndarray::Axis; use ndarray::parallel::Parallel; use ndarray::Dim; use ndarray::iter::AxisIterMut; use rayon::iter::IndexedParallelIterator; use ndarray::ViewRepr; use ndarray::ArrayBase; struct NdArrayMatrix { matrix: Array2<f32>, } impl NdArrayMatrix { pub fn new() -> Self { let matrix = Array2::zeros((10, 10)); Self { matrix, } } } trait ReadWrite { fn rw_read(&self, i: usize, j: usize) -> f32; fn rw_write(&mut self, i: usize, j: usize, val: f32); } impl ReadWrite for NdArrayMatrix { fn rw_read(&self, i: usize, j: usize) -> f32 { self.matrix[[i, j]] } fn rw_write(&mut self, i: usize, j: usize, val: f32) { self.matrix[[i, j]] = val; } } impl ReadWrite for ArrayBase<ViewRepr<&mut f32>, Dim<[usize; 1]>> { fn rw_read(&self, i: usize, j: usize) -> f32 { self[j] } fn rw_write(&mut self, i: usize, j: usize, val: f32) { self[j] = val; } } trait Sliceable<'a> { type IterItem: ReadWrite; type Output: IndexedParallelIterator<Item = Self::IterItem>; fn rows_par_iter(&'a mut self) -> Self::Output; } impl<'a> Sliceable<'a> for NdArrayMatrix { type IterItem = ArrayBase<ViewRepr<&'a mut f32>, Dim<[usize; 1]>>; type Output = Parallel<AxisIterMut<'a, f32, Dim<[usize; 1]>>>; fn rows_par_iter(&'a mut self) -> Self::Output { self.matrix.axis_iter_mut(Axis(0)).into_par_iter() } } fn main() { let mut matrix: NdArrayMatrix = NdArrayMatrix::new(); test(&mut matrix); } fn test<'a, T>(matrix: &'a mut T) where T: Sliceable<'a> + ReadWrite { matrix.rows_par_iter() .map(|mut row| { row.rw_write(0, 0, 0.0); }).count(); }
内容的提问来源于stack exchange,提问作者stomfaig
相关产品推荐
相关产品推荐

