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

如何为NumPy数组中的每个子数组设置不同起止索引进行切片?

为NumPy数组的每个子数组应用不同起止索引切片的高效方法

先明确你的场景:

import numpy as np
X = np.array([[1,2,3,4],[5,6,7,8]])

你需要实现类似自定义eg_slice函数的效果,传入每一行的起止索引对,输出对应切片后的二维数组,同时替代linspace生成全量索引的高内存方案。

以下是几种高效低内存的实现方式:

方法1:广播生成紧凑索引提取

这种方式避免生成全量冗余索引,仅按需创建必要的行、列索引数组:

def slice_per_row(X, slices):
    # 检查所有切片长度是否一致(否则无法组成二维数组)
    slice_lengths = [end - begin for begin, end in slices]
    if len(set(slice_lengths)) != 1:
        raise ValueError("所有切片的长度必须相同")
    length = slice_lengths[0]
    
    # 生成重复的行索引,对应每个切片元素
    row_idx = np.repeat(np.arange(X.shape[0]), length)
    # 拼接每行的连续列索引
    col_idx = np.concatenate([np.arange(begin, begin+length) for begin, _ in slices])
    
    return X[row_idx, col_idx].reshape(X.shape[0], length)

# 测试调用
result = slice_per_row(X, [[1,3],[0,2]])
print(result)
# 输出:
# [[2 3]
#  [5 6]]

方法2:利用广播直接索引(零额外内存开销)

通过广播机制生成索引视图,完全避免存储额外的索引数组,内存效率最高:

def slice_per_row_broadcast(X, slices):
    slice_lengths = [end - begin for begin, end in slices]
    if len(set(slice_lengths)) != 1:
        raise ValueError("所有切片的长度必须相同")
    length = slice_lengths[0]
    
    # 提取每行的起始索引,通过广播生成对应列索引矩阵
    starts = np.array([b for b, _ in slices])
    col_idx = starts[:, None] + np.arange(length)
    
    # 结合行索引广播完成切片
    return X[np.arange(X.shape[0])[:, None], col_idx]

# 测试调用
result = slice_per_row_broadcast(X, [[1,3],[0,2]])
print(result)
# 输出:
# [[2 3]
#  [5 6]]

方法3:列表推导式(简洁高效,适合小规模数组)

如果数组规模不大,直接循环处理每行,代码最简洁且无额外内存开销:

def slice_per_row_simple(X, slices):
    return np.array([X[i, begin:end] for i, (begin, end) in enumerate(slices)])

# 测试调用
result = slice_per_row_simple(X, [[1,3],[0,2]])
print(result)
# 输出:
# [[2 3]
#  [5 6]]

方案对比

上述所有方法都比linspace生成全量索引的方式更节省内存:linspace会生成完整的索引数组,而这些方法要么仅生成紧凑的必要索引,要么利用广播/视图机制避免存储冗余数据。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.08 11:15:58