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

如何在Rust中高效生成伯努利随机整数二维数组?

优化Rust中伯努利分布二维数组生成的性能

你的代码速度慢的核心原因是逐个元素采样+中间数组转换带来的额外开销,而NumPy的实现是基于批量优化的底层操作。下面是几种更高效的实现方式:

方法1:使用rand的批量填充方法

直接预分配u8数组,利用Rng::fill批量填充伯努利样本,避免中间布尔数组:

use rand::Rng;
use rand::distributions::Bernoulli;
use ndarray::Array2;

fn generate_bernoulli_array(nrows: usize, ncols: usize, p: f64) -> Array2<u8> {
    let mut rng = rand::thread_rng();
    let bernoulli = Bernoulli::new(p).unwrap();
    
    // 预分配全0数组
    let mut arr = Array2::<u8>::zeros((nrows, ncols));
    
    // 批量填充:Bernoulli会将成功(对应1)和失败(对应0)直接写入u8数组
    rng.fill(arr.as_slice_mut().unwrap(), &bernoulli);
    
    arr
}

方法2:使用ndarray-rand crate(最简洁高效)

ndarray-rand是ndarray官方的随机数扩展库,内部已经实现了批量生成的优化,代码更简洁:

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

ndarray = "0.15"
ndarray-rand = "0.14"
rand = "0.8"
rand_distr = "0.4"

然后编写代码:

use ndarray::Array2;
use ndarray_rand::rand_distr::Bernoulli;
use ndarray_rand::RandomExt;

fn generate_bernoulli_array(nrows: usize, ncols: usize, p: f64) -> Array2<u8> {
    // 直接生成符合伯努利分布的u8数组
    Array2::random((nrows, ncols), Bernoulli::new(p).unwrap())
}

原代码性能差的原因

  1. 逐个采样的开销:from_shape_fn遍历每个元素时,每次调用rng.sample(bernoulli)都会产生函数调用和分发的开销,大数组下这些开销会被放大。
  2. 中间布尔数组的冗余:先生成Array2<bool>再转换为u8数组,需要额外的一次全数组遍历,浪费内存带宽和CPU周期。
  3. 缺乏批量优化:rand的fill方法和ndarray-rand的实现都利用了批量采样的优化(比如减少随机数生成的重复操作、利用缓存局部性),比逐个采样高效得多。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.09 19:31:19