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

如何实现__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]

关键逻辑说明

  1. 映射表:_id_map存储字符串到整数的对应关系,可改为实例属性(比如初始化时传入)或动态修改,适配不同的索引映射需求。
  2. 单值索引处理:如果传入的是字符串,直接通过映射表转为整数;整数、None等numpy原生支持的索引类型直接返回。
  3. 切片处理:切片的start/stop/step可能包含字符串,需要递归调用处理函数,同时保留None(比如a[:]中的默认边界)。
  4. 多维度索引处理:当传入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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.20 19:05:39