Rust Ndarray crate中numpy roll()函数的替代实现方案是什么?
核心结论
截至ndarray 0.15.x最新稳定版本,官方没有内置与numpy roll 完全等价的公开接口,相关功能需求的issue目前仍处于开放状态,尚未合入正式发布版本。你可以通过非常轻量的代码自行实现该功能,逻辑和numpy原生roll完全一致,性能无额外损耗。
实现方案
1. 沿指定轴滚动(匹配你给出的示例行为)
你给出的示例预期输出,就是沿2维数组列轴(Axis(1))滚动1位的结果。numpy roll 传入axis参数时的行为是沿指定轴将数组元素循环位移,位移为正数时尾部的元素会移动到头部,位移为负数时头部元素移动到尾部。核心逻辑是沿目标轴将数组拆分为前后两段,交换顺序后拼接即可,实现代码如下:
use ndarray::{arr2, concatenate, Array, ArrayBase, Axis, Data, Dimension, Slice}; /// 沿指定axis对数组做循环滚动位移 /// # 参数 /// - arr: 输入数组 /// - shift: 位移量,正数向右/下滚动,负数向左/上滚动,自动对轴长度取模 /// - axis: 滚动的目标轴 fn roll<A, S, D>(arr: &ArrayBase<S, D>, shift: isize, axis: Axis) -> Array<A, D> where S: Data<Elem = A>, A: Clone, D: Dimension, { let axis_len = arr.len_of(axis); // 空数组直接返回 if axis_len == 0 { return arr.to_owned(); } // 位移量取模,兼容超过轴长度、负位移的场景 let valid_shift = shift.rem_euclid(axis_len as isize) as usize; if valid_shift == 0 { return arr.to_owned(); } // 拆分两段交换拼接 let tail_part = arr.slice_axis(axis, Slice::from((axis_len - valid_shift)..)); let head_part = arr.slice_axis(axis, Slice::from(..(axis_len - valid_shift))); concatenate(axis, &[tail_part.view(), head_part.view()]).unwrap() } // 测试示例场景 fn main() { let ar = arr2(&[[1.,2.,3.], [7., 8., 9.]]); // 沿第1轴(列方向)位移1位,和预期结果完全匹配 let rolled = roll(&ar, 1, Axis(1)); assert_eq!(rolled, arr2(&[[3.,1.,2.], [9.,7.,8.]])); }
2. 全局滚动(匹配numpy无axis参数的默认行为)
numpy调用roll不传入axis参数时,会先把数组展平为一维数组完成滚动,再reshape回原数组形状,实现代码如下:
use ndarray::{s, Axis}; fn roll_global<A, S, D>(arr: &ArrayBase<S, D>, shift: isize) -> Array<A, D> where S: Data<Elem = A>, A: Clone, D: Dimension, { let flat_arr = arr.flatten(); let total_len = flat_arr.len(); if total_len == 0 { return arr.to_owned(); } let valid_shift = shift.rem_euclid(total_len as isize) as usize; let rolled_flat = if valid_shift == 0 { flat_arr.to_owned() } else { let tail = flat_arr.slice(s![(total_len - valid_shift)..]); let head = flat_arr.slice(s![..(total_len - valid_shift)]); concatenate(Axis(0), &[tail, head]).unwrap() }; rolled_flat.into_shape_clone(arr.raw_dim()).unwrap() }
实现说明
- 上述实现的时间复杂度为O(n),和numpy原生roll性能一致:切片操作是零成本的视图操作,仅最终拼接时复制一次数据,无额外开销
- 自动处理位移值超过数组长度、负位移的边界场景,行为和numpy完全对齐
- 支持任意维度的ndarray数组,不局限于2维场景
如果不想自行维护实现代码,也可以使用ndarray生态中已实现roll功能的第三方扩展crate,注意选择和你使用的ndarray版本兼容的版本即可。
内容的提问来源于stack exchange,提问作者Anatoly Bugakov
相关产品推荐
相关产品推荐

