如何子类化np.ndarray实现右值转换、左值保留原始值?
子类化np.ndarray实现带转换逻辑的自定义数组
完全可行,通过重写numpy数组的核心方法,就能实现这种「运算用转换后值、赋值用原始值」的行为。以下是具体实现:
import numpy as np class MyNpArr(np.ndarray): def __new__(cls, input_array, transform_fwd): # 初始化自定义数组实例,绑定原始数据与转换函数 obj = np.asarray(input_array).view(cls) obj._raw_data = np.asarray(input_array).copy() obj._transform_fwd = transform_fwd return obj def __array_finalize__(self, obj): # 确保切片、视图等衍生实例继承原始数据与转换逻辑 if obj is None: return self._raw_data = getattr(obj, '_raw_data', None) self._transform_fwd = getattr(obj, '_transform_fwd', None) def __array_ufunc__(self, ufunc, method, *inputs, **kwargs): # 拦截所有numpy通用运算,将自定义数组转换为处理后的值参与计算 transformed_inputs = [] for inp in inputs: if isinstance(inp, MyNpArr): transformed_inputs.append(inp._transform_fwd(inp._raw_data)) else: transformed_inputs.append(inp) # 执行运算并返回普通numpy数组(匹配示例预期) result = ufunc(*transformed_inputs, **kwargs) return result.view(np.ndarray) if isinstance(result, np.ndarray) else result def __setitem__(self, key, value): # 赋值操作直接修改原始数据,同步更新数组视图 self._raw_data[key] = value super().__setitem__(key, value) def __getitem__(self, key): # 取值返回原始数据,切片返回带转换逻辑的自定义数组实例 raw_item = self._raw_data[key] return MyNpArr(raw_item, self._transform_fwd) if isinstance(raw_item, np.ndarray) else raw_item
测试代码(匹配示例逻辑)
# 使用log10转换函数,匹配示例输出 my_arr = MyNpArr([10.,100.,1000.], transform_fwd=lambda x: np.log10(x)) y = 2 * my_arr print(y) # 输出 [2. 4. 6.] my_arr[2] = 10000 y = 2 * my_arr print(y) # 输出 [2. 4. 8.]
关键逻辑说明
__new__:创建自定义数组实例时,保存原始数据副本与转换函数,同时让实例成为原始数据的视图。__array_ufunc__:处理所有numpy内置运算(如加减乘除、数学函数),确保运算时使用转换后的值,最终返回普通numpy数组。__setitem__:赋值操作直接修改原始数据,同步更新数组视图,保证后续运算能获取最新值。__array_finalize__:让切片、视图等衍生的自定义数组实例,自动继承原始数据与转换逻辑。
内容的提问来源于stack exchange,提问作者ipcamit
相关产品推荐
相关产品推荐

