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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 12:36:20