【问题标题】:Folding matrix rows element wise and computing average in Rust在 Rust 中按元素折叠矩阵行并计算平均值
【发布时间】:2021-04-04 02:43:52
【问题描述】:

假设我必须遵循Vec<Vec<_>> 表示的矩阵:

[[5, 9, 4],
 [8, 8, 2],
 [4, 5, 3]]

编译时未知的行数/列数,但保证行数等长。 我想明智地总结所有行元素并将它们除以行数(即,获得“平均列”)。最后我还需要得到Vec<_> 类型。 IE。以下:

[5.66, 7.33, 3] == [(5 + 8 + 4) / 3, (9 + 8 + 5) / 3, (4 + 2 + 3) / 3]

在 Rust 中最惯用的方法是什么?最好不要使用ndarray crate。

【问题讨论】:

  • 你是对的!对于那个很抱歉。修复了描述并用随机值更新以消除歧义

标签: vector rust iterator


【解决方案1】:

鉴于内部Vec 是一列,那么您可以将fold()sum() 交叉。

fn col_means(mat: &[Vec<f64>]) -> Vec<f64> {
    assert!(!mat.is_empty());

    let col_len = mat[0].len();
    let mut col_means = mat.iter().fold(vec![0.0; col_len], |mut col_means, row| {
        row.iter()
            .enumerate()
            .for_each(|(i, cell)| col_means[i] += cell);
        col_means
    });
    for col in col_means.iter_mut() {
        *col /= col_len as f64;
    }

    col_means
}

fn main() {
    let mat: Vec<Vec<f64>> = vec![
        vec![5.0, 9.0, 4.0],
        vec![8.0, 8.0, 2.0],
        vec![4.0, 5.0, 3.0],
    ];

    let col_means = col_means(&mat);
    println!("{:?}", col_means);
    // Outputs `[5.666666666666667, 7.333333333333333, 3.0]`
}

或者,您可以避免使用内部的Vec,然后使用chunks(),从而去除一层间接性。

fn col_means(mat: &[f64], col_len: usize) -> Vec<f64> {
    let mut col_means = mat
        .chunks(col_len)
        .fold(vec![0.0; col_len], |mut col_means, row| {
            row.iter()
                .enumerate()
                .for_each(|(i, cell)| col_means[i] += cell);
            col_means
        });
    for col in col_means.iter_mut() {
        *col /= col_len as f64;
    }

    col_means
}

fn main() {
    let mat: Vec<f64> = vec![5.0, 9.0, 4.0,
                             8.0, 8.0, 2.0,
                             4.0, 5.0, 3.0];

    let col_means = col_means(&mat, 3);
    println!("{:?}", col_means);
    // Outputs `[5.666666666666667, 7.333333333333333, 3.0]`
}

旧答案

假设内部Vec 是一行。然后您可以使用iter()sum()collect() 组合成Vec

let mat: Vec<Vec<i32>> = vec![vec![1, 2, 3], vec![2, 3, 4], vec![3, 4, 5]];

let row_means = mat
    .iter()
    .map(|row| row.iter().sum::<i32>() / (row.len() as i32))
    .collect::<Vec<_>>();

println!("{:?}", row_means);
// Outputs `[2, 3, 4]`

或者,您可以避免使用内部的Vec,然后使用chunks(),从而去除一层间接性。

let mat: Vec<i32> = vec![1, 2, 3, 2, 3, 4, 3, 4, 5];
let row_len = 3;

let row_means = mat
    .chunks(row_len as usize)
    .map(|row| row.iter().sum::<i32>() / row_len)
    .collect::<Vec<_>>();

println!("{:?}", row_means);
// Outputs `[2, 3, 4]`

然后您可以将实现抽象并隐藏到可重用的struct Mat

struct Mat(usize, Vec<i32>);

impl Mat {
    fn row_means(&self) -> Vec<i32> {
        let row_len = self.row_len();
        self.1
            .chunks(row_len)
            .map(|row| row.iter().sum::<i32>() / (row_len as i32))
            .collect()
    }

    fn cell(&self, row: usize, col: usize) -> i32 {
        self.1[col + row * self.row_len()]
    }

    fn row_len(&self) -> usize {
        self.0
    }

    fn col_len(&self) -> usize {
        self.1.len() / self.row_len()
    }
}

fn main() {
    let mat = Mat(3, vec![1, 2, 3,
                          2, 3, 4,
                          3, 4, 5]);
    println!("{:?}", mat.row_means());
    // Outputs `[2, 3, 4]`
}

【讨论】:

  • 确实搞砸了第一个。我现在已经修好了。我通常认为x_{count,size,len} 是同义词,但是是的,在这种情况下,我可以看到row_count 听起来是错误的。与row_size 相比,我已将其重命名为row_len 与Rust 更接近:)
【解决方案2】:

我写了一个函数来计算任意大小网格的平均列:

#[derive(Debug)]
struct EmptyGrid;

fn mean_column(grid: &[Vec<f64>]) -> Result<Vec<f64>, EmptyGrid> {
    if grid.is_empty() {
        return Err(EmptyGrid);
    }
    if grid[0].is_empty() {
        return Ok(Vec::new());
    }
    let mut mean = grid[0].to_vec();
    let col_len = mean.len();
    for row in grid.iter().skip(1) {
        for i in 0..col_len {
            mean[i] += row[i];
        }
    }
    let row_len = grid.len() as f64;
    for sum in mean.iter_mut() {
        *sum /= row_len;
    }
    Ok(mean)
}

fn main() {
    let grid = vec![
        vec![5.0, 9.0, 4.0],
        vec![8.0, 8.0, 2.0],
        vec![4.0, 5.0, 3.0],
    ];
    let mean_column = mean_column(&grid).unwrap();
    dbg!(mean_column); // prints [5.66, 7.33, 3.0]
}

playground

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2015-04-13
    • 1970-01-01
    • 2021-07-24
    • 2013-05-06
    • 1970-01-01
    相关资源
    最近更新 更多