如何在Rust ndarray中实现类似NumPy的切片赋值操作?
Rust ndarray 多维数组切片赋值实现
背景
在Python的NumPy中,我们可以对多维数组的切片直接完成赋值操作,示例如下:
import numpy as np a = np.ones((4, 5, 6, 7)) b = np.zeros((4, 5, 6, 7)) # 将b的指定切片赋值给a的对应切片 a[1:3, 0:2, 1::2, :] = b[1:3, 0:2, 1::2, :]
现在有两个Rust ndarray的4维数组:
use ndarray::prelude::*; fn main() { let mut a = Array4::<f64>::ones((4, 5, 6, 7)); let b = Array4::<f64>::zeros((4, 5, 6, 7)); }
需要实现和上述NumPy完全等价的切片赋值逻辑。
实现方案
在Rust ndarray中,有两种常用方式完成该操作,核心是利用切片索引结合赋值方法:
方法一:使用切片宏与.assign()方法
这是最简洁的实现方式,语法和NumPy切片高度对齐:
use ndarray::prelude::*; fn main() { let mut a = Array4::<f64>::ones((4, 5, 6, 7)); let b = Array4::<f64>::zeros((4, 5, 6, 7)); // 对a取可变切片,对b取不可变切片,调用assign完成赋值 a.slice_mut(s![1..3, 0..2, 1..;2, ..]) .assign(&b.slice(s![1..3, 0..2, 1..;2, ..])); }
切片宏s![]的语法对应关系:
1..3对应NumPy的1:3(左闭右开区间)0..2对应0:21..;2对应1::2(起始索引为1,步长为2)..对应:(取当前维度的全部元素)
方法二:迭代器遍历赋值
如果需要逐个元素的操作灵活性,也可以通过迭代器遍历两个切片的元素完成赋值:
use ndarray::prelude::*; fn main() { let mut a = Array4::<f64>::ones((4, 5, 6, 7)); let b = Array4::<f64>::zeros((4, 5, 6, 7)); let mut a_slice = a.slice_mut(s![1..3, 0..2, 1..;2, ..]); let b_slice = b.slice(s![1..3, 0..2, 1..;2, ..]); // 遍历两个切片的元素对,完成赋值 for (a_elem, b_elem) in a_slice.iter_mut().zip(b_slice.iter()) { *a_elem = *b_elem; } }
注意事项
- 必须保证两个切片的形状完全一致,否则
.assign()会触发panic; - 操作
a的可变切片时,a必须是可变绑定(即mut a); - 切片宏
s![]需要引入ndarray::prelude::*才能使用。
内容的提问来源于stack exchange,提问作者Matt Pitkin
相关产品推荐
相关产品推荐

