如何从numpy数组切片生成的视图中还原对应的索引?
实现方案
这个需求完全可实现,不需要反向推导切片参数,直接在切片操作触发时捕获索引参数,同步应用到自定义属性上即可,实现逻辑如下:
核心思路
NumPy的切片操作本质是调用数组的__getitem__方法,传入的参数就是用户写的索引规则。我们只需要重写子类的__getitem__方法,在生成切片视图的同时,把相同索引应用到你的自定义属性数组上,就能实现属性同步切片。
完整实现代码
import numpy as np class CustomArray(np.ndarray): def __new__(cls, input_array, attr_array=None): # 生成自定义类的数组实例 obj = np.asarray(input_array).view(cls) # 校验自定义属性的维度匹配 if attr_array is not None: attr_array = np.asarray(attr_array) if attr_array.shape != obj.shape: raise ValueError("属性数组维度需和主数组一致") obj.attr = attr_array return obj def __array_finalize__(self, obj): # 处理视图/拷贝生成时的属性初始化,避免属性丢失 if obj is None: return self.attr = getattr(obj, 'attr', None) def __getitem__(self, key): # 调用父类方法生成切片后的视图/实例 sliced_main = super().__getitem__(key) # 仅当返回的是CustomArray实例(非单个元素标量)时同步属性 if isinstance(sliced_main, CustomArray): sliced_main.attr = self.attr[key] return sliced_main
使用示例
# 初始化主数组和同维度属性数组 a = CustomArray(np.arange(10), attr_array=np.arange(10)*10) # 执行切片操作 b = a[:5] # 验证结果 print(b) # 输出 [0 1 2 3 4] print(b.attr) # 输出 [ 0 10 20 30 40],属性数组同步完成切片
补充说明
- 该方案天然支持所有NumPy索引规则:普通切片、步长切片、布尔索引、整数数组索引、多维索引都可以正常适配,不需要额外解析索引参数
- 如果需要处理ufunc运算、原位修改等场景下的属性同步,可以额外重写
__array_ufunc__方法扩展逻辑,仅做切片同步的话上述代码即可满足需求 - 反向推导切片参数的方案可靠性极低,遇到非连续内存、组合索引等场景很容易出错,不推荐使用
内容的提问来源于stack exchange,提问作者cymin
相关产品推荐
相关产品推荐

