You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Rust:实现扁平化存储矩阵的可变列迭代器的借用检查问题

问题:行优先矩阵的可变列迭代器实现

我有一个以**行优先(row-major)**方式存储在Vec<T>中的矩阵,目前用嵌套循环遍历列元素:

for i in 0..width {
    for j in 0..height {
        let element = matrix.get_mut(i, j);
        // do something
    }
}

这种写法可行,但我想用迭代器优化代码 ergonomics,期望实现一个返回类型为fn cols_mut() -> impl ExactSizeIterator<Item = impl ExactSizeIterator<Item = &mut T>>的函数,这样就能像下面这样使用:

matrix.cols_mut().for_each(|col| {
    col.fold(0, |sum, val| {
        *val += sum;
        *val
    });
});

不可变版本的cols()实现很简单:

fn cols(&self) -> impl ExactSizeIterator<Item = impl ExactSizeIterator<Item = &T>> {
    (0..self.width()).map(move |i| {
        (0..self.height()).map(move |j| {
            self.get(i, j).unwrap()
        })
    })
}

但可变版本始终无法通过Rust借用检查器,求解决方法。


解决方案

核心问题在于:Rust的借用检查器不允许同时存在多个指向同一数据的可变借用,直接模仿不可变版本的写法会尝试多次获取整个矩阵的可变引用,触发检查错误。下面提供两种可行的实现方式:

方式一:使用unsafe代码(安全可控)

因为我们明确知道列迭代器中的每个元素都是矩阵中唯一的位置,不会出现重叠借用,所以可以用unsafe绕过借用检查器的限制:

首先定义基础矩阵结构体(假设你的结构体结构如下):

struct Matrix<T> {
    data: Vec<T>,
    width: usize,
    height: usize,
}

impl<T> Matrix<T> {
    fn new(width: usize, height: usize, data: Vec<T>) -> Self {
        assert_eq!(width * height, data.len());
        Matrix { data, width, height }
    }

    fn width(&self) -> usize {
        self.width
    }

    fn height(&self) -> usize {
        self.height
    }

    fn get(&self, col: usize, row: usize) -> Option<&T> {
        if col >= self.width || row >= self.height {
            None
        } else {
            self.data.get(row * self.width + col)
        }
    }
}

实现cols_mut函数:

impl<T> Matrix<T> {
    fn cols_mut(&mut self) -> impl ExactSizeIterator<Item = impl ExactSizeIterator<Item = &mut T>> {
        (0..self.width).map(move |col_idx| {
            let data = &mut self.data;
            let width = self.width;
            (0..self.height).map(move |row_idx| {
                let idx = row_idx * width + col_idx;
                // 安全:每个索引只会被访问一次,无重叠可变借用
                unsafe { &mut *data.as_mut_ptr().add(idx) }
            })
        })
    }
}

这里的unsafe是安全的,每个矩阵元素只会被一个迭代器访问一次,完全符合Rust内存安全规则,只是借用检查器无法自动推断这一点。

方式二:自定义迭代器结构体(零unsafe,符合Rust风格)

如果不想使用unsafe,可以通过自定义迭代器结构体跟踪位置,用split_mut确保每次只返回唯一的可变引用:

先定义两个迭代器结构体:

struct ColMutIter<'a, T> {
    matrix: &'a mut Matrix<T>,
    current_col: usize,
}

struct ColElementMutIter<'a, T> {
    data: &'a mut Vec<T>,
    width: usize,
    col: usize,
    current_row: usize,
}

// 实现列迭代器的Iterator和ExactSizeIterator trait
impl<'a, T> Iterator for ColMutIter<'a, T> {
    type Item = ColElementMutIter<'a, T>;

    fn next(&mut self) -> Option<Self::Item> {
        if self.current_col >= self.matrix.width {
            None
        } else {
            let col = self.current_col;
            self.current_col += 1;
            Some(ColElementMutIter {
                data: &mut self.matrix.data,
                width: self.matrix.width,
                col,
                current_row: 0,
            })
        }
    }

    fn size_hint(&self) -> (usize, Option<usize>) {
        let remaining = self.matrix.width - self.current_col;
        (remaining, Some(remaining))
    }
}

impl<'a, T> ExactSizeIterator for ColMutIter<'a, T> {}

// 实现列元素迭代器的Iterator和ExactSizeIterator trait
impl<'a, T> Iterator for ColElementMutIter<'a, T> {
    type Item = &'a mut T;

    fn next(&mut self) -> Option<Self::Item> {
        let height = self.data.len() / self.width;
        if self.current_row >= height {
            None
        } else {
            let idx = self.current_row * self.width + self.col;
            self.current_row += 1;
            // 通过split_mut拆分数据,确保每次只返回一个唯一的可变引用
            let (front, back) = self.data.split_mut_at(idx + 1);
            Some(&mut front[idx])
        }
    }

    fn size_hint(&self) -> (usize, Option<usize>) {
        let height = self.data.len() / self.width;
        let remaining = height - self.current_row;
        (remaining, Some(remaining))
    }
}

impl<'a, T> ExactSizeIterator for ColElementMutIter<'a, T> {}

最后在Matrix结构体中实现cols_mut:

impl<T> Matrix<T> {
    fn cols_mut(&mut self) -> ColMutIter<'_, T> {
        ColMutIter {
            matrix: self,
            current_col: 0,
        }
    }
}

测试示例

fn main() {
    let mut matrix = Matrix::new(3, 2, vec![1, 2, 3, 4, 5, 6]);
    matrix.cols_mut().for_each(|mut col| {
        let mut sum = 0;
        col.for_each(|val| {
            sum += *val;
            *val = sum;
        });
    });
    // 输出修改后的矩阵:[1, 3, 6, 4, 9, 15]
    println!("{:?}", matrix.data);
}

内容的提问来源于stack exchange,提问作者Ymi_Yugy

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.20 11:22:55