如何在自定义数组类中实现numpy风格的索引功能?
实现自定义numpy兼容数组索引的优化方案
优先推荐:复用numpy原生索引能力,完全避免手写判断
- 直接调用numpy公开的索引解析API
np.lib.index_tricks.normalize_index,传入你的数组shape和接收到的key参数,即可直接输出符合numpy规范的标准化索引元组,自动覆盖整数、切片、布尔数组、花式索引、多维度混合索引、np.newaxis等所有合法索引场景,不需要自己写任何类型判断逻辑。 - 拿到标准化索引后,你可以直接将其应用到你类内部存储的实际数据结构上,如果内部本身就维护了一个numpy数组作为实际存储,直接用这个索引操作内部数组即可,完美匹配numpy的行为。
更省心的兼容方案:继承numpy原生类
- 如果你的自定义数组不需要完全从零实现底层存储,直接继承
np.ndarray是成本最低的方案,numpy会默认帮你实现所有索引、广播、通用函数的逻辑,天然就是numpy兼容的drop-in替换,你只需要重写你需要自定义逻辑的特定方法即可,完全不需要自行维护__getitem__和__setitem__的实现。 - 如果你是要包裹其他非numpy的底层存储结构,可以实现numpy的
__array_function__和__array_ufunc__协议,仅拦截你需要自定义处理的操作,剩下的所有逻辑包括索引解析全部交由numpy默认处理,也能大幅减少自行维护的代码量。
必须从零实现索引逻辑的优化方案
如果因为特殊限制必须自己手写索引逻辑,按以下方式优化维护成本:
- 把索引解析、合法性校验、广播匹配的逻辑抽成独立的公共工具函数,
__getitem__和__setitem__复用同一套逻辑,避免两份代码重复维护。 - 广播规则不需要自己手写,直接调用
np.broadcast_shapes判断索引输出形状和value的形状是否兼容即可。
内容的提问来源于stack exchange,提问作者Bubaya
相关产品推荐
相关产品推荐

