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

Rust并行矩阵乘法性能劣于串行的原因排查

Rust并行矩阵乘法性能劣于串行的原因分析

问题背景

我在Rust中实现了串行与并行矩阵乘法,并行版本按CPU核心数拆分任务到独立线程,但100x100矩阵的基准测试显示并行版本耗时是串行的3倍左右:

Time elapsed in sequential multiplication is: 2.71852745s
Time elapsed in parallel multiplication is: 8.5174711s

核心并行实现代码片段:

fn mul_par(&self, other: &Self) -> Result<Self, String> {
    if self.cols != other.rows {
            return Err(format!(
                "Invalid matrix dimensions. \
                The number of columns of the first matrix should be equal to the number of rows \
                of the second matrix. Got {} and {} instead.",
            self.cols, other.rows,
        ));
    }

    use num_cpus;
    let num_threads = num_cpus::get();
    let result = Arc::new(Mutex::new(vec![vec![0.0; other.cols]; self.rows]));
    let chunk_size: usize = (self.rows as f32 / num_threads as f32).ceil() as usize;
    let mut handles = Vec::with_capacity(num_threads);

    for th in 0..num_threads {
        let start = th * chunk_size;
        if start >= self.rows {
            break;
        };
        let result = Arc::clone(&result);
        let self_data = Arc::clone(&self.data);
        let other_data = Arc::clone(&other.data);

        let is_last_thread = th == num_threads - 1;
        let end = if is_last_thread {
            self.rows
        } else {
            (th + 1) * chunk_size
        };

        let other_cols = other.cols;
        let self_cols = self.cols;

        let handle = thread::spawn(move || {
            for row_a in start..end {
                for col_b in 0..other_cols {
                    let mut sum = 0.0;
                    for i in 0..self_cols {
                        sum += self_data[row_a][i] * other_data[i][col_b];
                    }
                    // 每个元素计算后都加锁写入
                    result.lock().unwrap()[row_a][col_b] = sum;
                }
            }
        });
        handles.push(handle);
    }

    for handle in handles {
        handle.join().unwrap();
    }

    let result = Arc::try_unwrap(result).unwrap().into_inner().unwrap();
    Ok(Self::new(result)?)
}

核心性能问题原因

1. 全局Mutex的频繁竞争

并行代码中,每个元素计算完成后都调用result.lock().unwrap()写入结果,这会导致所有线程频繁争抢同一个全局锁。锁的竞争会引发大量上下文切换,其开销远远超过并行计算带来的收益,直接导致并行版本效率暴跌。

2. 内存布局的缓存不友好

当前矩阵使用Vec<Vec<f64>>(嵌套向量)存储,属于行优先的非连续内存布局。在矩阵乘法的内层循环中,访问other_data[i][col_b]是按列读取数据,这会导致CPU缓存命中率极低——因为内存中连续存储的是行数据,列访问会跳过大量内存块,频繁触发缓存失效,严重拖慢计算速度。串行版本已经受此影响,并行版本多线程的缓存竞争会进一步加剧这个问题。

3. 小矩阵的并行开销过高

对于100x100的小矩阵,计算总量本身很小,但并行版本需要创建多个线程、克隆Arc指针,这些操作的固定开销占比极高,完全抵消了并行计算的优势。

优化方案

1. 消除锁竞争:每个线程独立生成结果块

让每个线程计算自己负责的行的所有结果,存储在局部向量中,最后再合并所有线程的结果,完全避免全局锁:

fn mul_par(&self, other: &Self) -> Result<Self, String> {
    if self.cols != other.rows {
        return Err(format!(
            "Invalid matrix dimensions. Got {} cols vs {} rows",
            self.cols, other.rows
        ));
    }

    use num_cpus;
    let num_threads = num_cpus::get();
    let chunk_size = (self.rows + num_threads - 1) / num_threads; // 向上取整
    let mut handles = Vec::with_capacity(num_threads);

    for th in 0..num_threads {
        let start = th * chunk_size;
        let end = (start + chunk_size).min(self.rows);
        if start >= end {
            break;
        }

        let self_data = Arc::clone(&self.data);
        let other_data = Arc::clone(&other.data);
        let other_cols = other.cols;
        let self_cols = self.cols;

        handles.push(thread::spawn(move || {
            let mut local_result = vec![vec![0.0; other_cols]; end - start];
            for (local_row, row_a) in (start..end).enumerate() {
                for col_b in 0..other_cols {
                    let mut sum = 0.0;
                    for i in 0..self_cols {
                        sum += self_data[row_a][i] * other_data[i][col_b];
                    }
                    local_result[local_row][col_b] = sum;
                }
            }
            local_result
        }));
    }

    // 合并所有线程的局部结果
    let mut result = Vec::with_capacity(self.rows);
    for handle in handles {
        result.extend(handle.join().unwrap());
    }

    Ok(Self::new(result)?)
}

2. 优化内存布局:使用连续一维数组

将矩阵改为连续的一维数组存储,通过row * cols + col计算索引,同时调整循环顺序提升缓存命中率:

#[derive(Debug)]
struct RsMatrix {
    data: Arc<Vec<f64>>,
    rows: usize,
    cols: usize,
}

impl RsMatrix {
    fn new(data: Vec<Vec<f64>>) -> Result<Self, String> {
        let row_len = data[0].len();
        if !data.iter().all(|row| row.len() == row_len) {
            return Err(String::from("All rows must have the same length"));
        }
        let rows = data.len();
        let cols = row_len;
        // 转换为连续一维数组
        let flat_data = data.into_iter().flatten().collect();
        Ok(Self { data: Arc::new(flat_data), rows, cols })
    }

    // 获取元素的辅助方法
    fn get(&self, row: usize, col: usize) -> f64 {
        self.data[row * self.cols + col]
    }
}

同时调整乘法循环顺序为row_a -> i -> col_b,让内存访问更连续:

// 并行版本中的计算循环
for (local_row, row_a) in (start..end).enumerate() {
    for i in 0..self_cols {
        let val_a = self_data.get(row_a, i);
        for col_b in 0..other_cols {
            local_result[local_row][col_b] += val_a * other_data.get(i, col_b);
        }
    }
}

3. 使用高效并行库Rayon

Rayon提供了简单的并行迭代器,自动处理任务拆分和线程管理,比手动创建线程更高效:

// 添加依赖:rayon = "1.8"
use rayon::prelude::*;

fn mul_par_rayon(&self, other: &Self) -> Result<Self, String> {
    if self.cols != other.rows {
        return Err(format!(
            "Invalid matrix dimensions. Got {} cols vs {} rows",
            self.cols, other.rows
        ));
    }

    let result_data = (0..self.rows)
        .into_par_iter()
        .map(|row_a| {
            let mut row = vec![0.0; other.cols];
            for col_b in 0..other.cols {
                let mut sum = 0.0;
                for i in 0..self.cols {
                    sum += self.get(row_a, i) * other.get(i, col_b);
                }
                row[col_b] = sum;
            }
            row
        })
        .collect();

    Ok(Self::new(result_data)?)
}

4. 设置并行阈值

只在矩阵足够大时启用并行,避免小矩阵的并行开销:

fn mul_par(&self, other: &Self) -> Result<Self, String> {
    // 当行数小于阈值时,直接使用串行版本
    if self.rows < 500 {
        return self.mul(other);
    }
    // 并行逻辑...
}

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 21:02:11