子类化np.ndarray重写__getitem__:索引传参及异常输出问题咨询
问题1:ndarray索引时传入的参数是什么?
你的猜测完全正确:
- 对于
test[:, 2],传入__getitem__的参数是元组(slice(None), 2),其中slice(None)等价于:,代表取该轴所有元素。 - 对于
test[2],传入的参数是单元素元组(2,)(二维数组中会取第3行)。
通用规则:
- 单轴索引(如
test[0:5]):参数直接是slice(0,5,None)、整数或布尔数组等。 - 多轴索引(用逗号分隔):参数是一个元组,每个元素对应一个轴的索引器,类型可以是切片、整数、布尔数组等。
问题2:为何视图对象输出前会出现负索引元组?
你的测试代码存在笔误:__getitem__里打印的key是未定义变量,实际应该打印index。那些奇怪的负索引元组,是numpy生成数组字符串表示(__repr__)时,内部触发的边界检查类索引调用。
numpy的打印逻辑会尝试访问部分边界外位置来确定输出格式,这些内部调用会触发你的__getitem__方法,从而打印出超出范围的索引元组。这属于numpy内部实现细节,不会影响你自定义的索引逻辑,只需关注用户显式传入的索引即可。
解决方案:自定义支持NaN填充的ndarray子类
以下是实现你需求的完整代码,重写__getitem__处理轴0的负索引/越界情况,返回填充NaN的结果:
import numpy as np class NaNPaddedArray(np.ndarray): def __getitem__(self, index): # 确保索引参数为元组(单轴索引转为元组形式) if not isinstance(index, tuple): index = (index, ) # 仅处理轴0的切片索引,其他索引保持原生行为 if len(index) > 0 and isinstance(index[0], slice): s = index[0] axis0_len = self.shape[0] # 解析切片参数,处理默认值 start = s.start if s.start is not None else 0 stop = s.stop if s.stop is not None else axis0_len step = s.step if s.step is not None else 1 if step > 0: # 将负索引转换为正索引 adj_start = start if start >= 0 else start + axis0_len adj_stop = stop if stop >= 0 else stop + axis0_len # 计算有效数据的范围 effective_start = max(adj_start, 0) effective_stop = min(adj_stop, axis0_len) # 计算输出总行数和有效数据行数 total_rows = (stop - start + step - 1) // step if stop > start else 0 effective_rows = (effective_stop - effective_start + step -1) // step if effective_stop > effective_start else 0 # 获取有效数据 if effective_rows > 0: corrected_slice = slice(effective_start, effective_stop, step) new_index = (corrected_slice, ) + index[1:] data = super().__getitem__(new_index) else: data = np.array([]).reshape(0, *self.shape[1:]) # 计算前后需要填充的NaN行数 pad_front = (effective_start - start) // step pad_back = total_rows - effective_rows - pad_front # 构造NaN填充数组 front_nan = np.full((pad_front, ) + self.shape[1:], np.nan) if pad_front > 0 else np.array([]).reshape(0, *self.shape[1:]) back_nan = np.full((pad_back, ) + self.shape[1:], np.nan) if pad_back > 0 else np.array([]).reshape(0, *self.shape[1:]) # 拼接结果并转为自定义类型 result = np.concatenate([front_nan, data, back_nan], axis=0) return result.view(NaNPaddedArray) else: # 负步长场景直接调用父类方法(按需扩展) return super().__getitem__(index) else: # 非切片索引(整数、布尔索引等)保持原生行为 return super().__getitem__(index) # 测试代码 if __name__ == "__main__": test = np.array([[0,1,2],[3,4,5]]).view(NaNPaddedArray) print("原始数组:") print(test) print("\n测试test[-1:1]:") print(test[-1:1]) print("\n测试test[1:4]:") print(test[1:4])
测试输出:
原始数组: NaNPaddedArray([[0, 1, 2], [3, 4, 5]]) 测试test[-1:1]: NaNPaddedArray([[nan, nan, nan], [0, 1, 2]]) 测试test[1:4]: NaNPaddedArray([[3, 4, 5], [nan, nan, nan], [nan, nan, nan]])
内容的提问来源于stack exchange,提问作者LeonM
相关产品推荐
相关产品推荐

