如何实现__getitem__以预处理参数后传入底层numpy数组?
实现支持自定义索引映射的Numpy包装类
要解决这个问题,核心是对__getitem__接收的所有索引类型(单值、多维度元组、切片)做递归预处理,将需要映射的字符串转换为numpy能识别的整数,再传递给内部的self._mat。以下是完整实现:
import numpy as np class A(): # 定义字符串到整数的映射表,可按需修改或扩展 _id_map = {'deadbeef': 0} def __init__(self, n: int=3): self._mat = np.zeros(shape=[n, n]) def _process_val(self, val): # 处理字符串索引:通过映射表转换为整数 if isinstance(val, str): return self._id_map[val] # 处理切片对象:递归处理切片的start/stop/step elif isinstance(val, slice): return slice( self._process_val(val.start) if val.start is not None else None, self._process_val(val.stop) if val.stop is not None else None, self._process_val(val.step) if val.step is not None else None ) # 处理多维度元组索引:逐个处理元组内的每个元素 elif isinstance(val, tuple): return tuple(self._process_val(item) for item in val) # 其他合法索引类型(整数、None、Ellipsis等)直接返回 else: return val def __getitem__(self, val): processed_idx = self._process_val(val) return self._mat[processed_idx]
关键逻辑说明
- 映射表:
_id_map存储字符串到整数的对应关系,可改为实例属性(比如初始化时传入)或动态修改,适配不同的索引映射需求。 - 单值索引处理:如果传入的是字符串,直接通过映射表转为整数;整数、
None等numpy原生支持的索引类型直接返回。 - 切片处理:切片的
start/stop/step可能包含字符串,需要递归调用处理函数,同时保留None(比如a[:]中的默认边界)。 - 多维度索引处理:当传入
a[0,1]或a[(0,1)]时,实际接收的是元组参数,需要遍历元组内每个元素逐一处理,返回新的元组作为索引。
示例验证
运行以下代码可验证效果:
a = A(n=4) # 常规索引测试 print(a[0]) # 输出:[0. 0. 0. 0.] print(a[0, 1]) # 输出:0.0 print(a[(0, 1)]) # 输出:0.0 print(a[:, 1:2]) # 输出:[[0.] # [0.] # [0.] # [0.]] # 字符串索引测试 print(a['deadbeef']) # 输出:[0. 0. 0. 0.](等价于a[0]) print(a['deadbeef', 1]) # 输出:0.0(等价于a[0,1]) print(a[('deadbeef', 1)]) # 输出:0.0(等价于a[(0,1)]) print(a[:, 'deadbeef':2]) # 输出:[[0. 0.] # [0. 0.] # [0. 0.] # [0. 0.]](等价于a[:,0:2])
扩展优化
如果需要支持字符串数组这类高级索引,可以在_process_val中添加数组处理逻辑:
elif isinstance(val, np.ndarray) and val.dtype.kind in ('U', 'S'): # 批量转换字符串数组为整数数组 return np.vectorize(lambda x: self._id_map[x])(val)
内容的提问来源于stack exchange,提问作者Tal Afek
相关产品推荐
相关产品推荐

