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

Rust中无互斥实现多线程数组外积计算的正确方法

问题描述

我尝试实现一个计算两个一维数组外积的outer函数,编写了多线程版本的代码。设计上各线程不会写入同一元素,但Rust编译器不允许同一result变量存在多个可变引用,而使用互斥锁会大幅降低性能,请问该如何正确实现此函数?

原代码如下:

use std::thread;
use ndarray::prelude::*;

pub fn multithread_outer(A: &Array1<f64>, B: &Array1<f64>) -> Array2<f64> {
    let mut result = Array2::<f64>::default((A.len(), B.len()));
    let thread_num = 5;
    let n = A.len() / thread_num;

    // a & b are ArcArray2<f64>
    let a = A.to_owned().into_shared();
    let b = B.to_owned().into_shared();

    for i in 0..thread_num{
        let a = a.clone();
        let b = b.clone();
        thread::spawn(move || {
            for j in i * n..(i + 1) * n {
                for k in 0..b.len() {
                    // This is the line I want to change
                    result[[j, k]] = a[j] * b[k];
                }
            }
        });
    }
    // Use join to make sure all threads finish here
    // Not so related to this question, so I didn't put it here
    result
}
解决方案

针对这个问题,有几种安全且高效的实现方式,不需要使用互斥锁:

方法一:手动拆分可变切片

利用Rust的借用规则,将结果数组拆分为互不重叠的可变切片,每个线程只操作自己负责的切片区域。因为各线程处理的是不同的行区间,不存在数据竞争,编译器会认可这种安全的拆分。

修改后的代码如下:

use std::thread;
use ndarray::prelude::*;

pub fn multithread_outer(A: &Array1<f64>, B: &Array1<f64>) -> Array2<f64> {
    let mut result = Array2::<f64>::default((A.len(), B.len()));
    let thread_num = 5;
    let n = A.len() / thread_num;
    let remainder = A.len() % thread_num;

    let a = A.to_owned().into_shared();
    let b = B.to_owned().into_shared();

    let mut handles = Vec::with_capacity(thread_num);

    // 按行拆分结果数组为多个可变切片
    let mut result_chunks = result.axis_chunks_mut(Axis(0), n);
    for (i, mut chunk) in result_chunks.into_iter().enumerate() {
        let a = a.clone();
        let b = b.clone();
        // 处理最后一个线程的剩余行
        let start = i * n;
        let end = if i == thread_num - 1 { start + n + remainder } else { start + n };

        let handle = thread::spawn(move || {
            for j in start..end {
                // 获取当前行的可变引用,j-start是切片内的相对索引
                let row = chunk.row_mut(j - start);
                let a_val = a[j];
                // 填充当前行的所有元素
                row.iter_mut().zip(B.iter()).for_each(|(val, &b_val)| {
                    *val = a_val * b_val;
                });
            }
        });
        handles.push(handle);
    }

    // 等待所有线程完成计算
    for handle in handles {
        handle.join().unwrap();
    }

    result
}

方法二:使用ndarray的并行迭代(推荐)

ndarray库支持通过rayon特性实现并行迭代,这种方式更简洁,不需要手动管理线程和切片拆分,底层会自动安全地分配任务给多个线程,性能也有保障。

首先需要在Cargo.toml中添加依赖:

ndarray = { version = "0.15", features = ["rayon"] }

实现代码如下:

use ndarray::prelude::*;
use ndarray::parallel::prelude::*;

pub fn parallel_outer(A: &Array1<f64>, B: &Array1<f64>) -> Array2<f64> {
    let mut result = Array2::<f64>::default((A.len(), B.len()));
    
    // 并行遍历结果数组的每一行,与A的对应元素配对
    result.axis_iter_mut(Axis(0))
          .into_par_iter()
          .zip(A.iter())
          .for_each(|(mut row, &a_val)| {
              // 填充当前行的所有元素
              row.iter_mut().zip(B.iter()).for_each(|(val, &b_val)| {
                  *val = a_val * b_val;
              });
          });
    
    result
}

方法三:使用unsafe(不推荐)

如果一定要手动控制内存访问,可以用UnsafeCell配合Arc绕开借用检查,但这种方法需要你自己保证绝对没有数据竞争,一旦逻辑出错会导致未定义行为,因此仅作为兜底方案:

use std::sync::{Arc, UnsafeCell};
use std::thread;
use ndarray::prelude::*;

pub fn unsafe_multithread_outer(A: &Array1<f64>, B: &Array1<f64>) -> Array2<f64> {
    let result = Arc::new(UnsafeCell::new(Array2::<f64>::default((A.len(), B.len()))));
    let thread_num = 5;
    let n = A.len() / thread_num;
    let remainder = A.len() % thread_num;

    let a = A.to_owned().into_shared();
    let b = B.to_owned().into_shared();

    let mut handles = Vec::with_capacity(thread_num);

    for i in 0..thread_num {
        let result = Arc::clone(&result);
        let a = a.clone();
        let b = b.clone();
        let start = i * n;
        let end = if i == thread_num - 1 { start + n + remainder } else { start + n };

        let handle = thread::spawn(move || {
            // 手动获取可变引用,需保证无数据竞争
            let result = unsafe { &mut *result.get() };
            for j in start..end {
                for k in 0..b.len() {
                    result[[j, k]] = a[j] * b[k];
                }
            }
        });
        handles.push(handle);
    }

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

    // 取出最终结果
    Arc::try_unwrap(result).unwrap().into_inner()
}

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 20:46:34