如何在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()) }
原代码性能差的原因
- 逐个采样的开销:
from_shape_fn遍历每个元素时,每次调用rng.sample(bernoulli)都会产生函数调用和分发的开销,大数组下这些开销会被放大。 - 中间布尔数组的冗余:先生成
Array2<bool>再转换为u8数组,需要额外的一次全数组遍历,浪费内存带宽和CPU周期。 - 缺乏批量优化:rand的
fill方法和ndarray-rand的实现都利用了批量采样的优化(比如减少随机数生成的重复操作、利用缓存局部性),比逐个采样高效得多。
内容的提问来源于stack exchange,提问作者Abdullah Khalid
相关产品推荐
相关产品推荐

