如何为自定义数组容器实现np.where?解决递归错误问题
自定义数组容器适配np.where时的递归错误问题
我尝试在自定义类myClass中通过myClass.data属性使用np.ndarray,参考NumPy官方文档《编写自定义数组容器》开发。已成功为np.sum、np.mean、np.std注册array_function实现,但适配np.where时触发递归错误。
示例代码
from numpy.typing import ArrayLike import numpy as np import pandas as pd import numpy.lib.mixins from numbers import Number HANDLED_FUNCTIONS = {} def implements(np_function): def decorator(func): HANDLED_FUNCTIONS[np_function] = func return func return decorator class myClass(numpy.lib.mixins.NDArrayOperatorsMixin): def __init__(self, name, data: ArrayLike, index: ArrayLike): self.name = name self.data = data self.index = index def __array__(self, dtype=None, copy=None): if copy is False: raise ValueError( "`copy=False` isn't supported. A copy is always created." ) return self.data def __array_ufunc__(self, ufunc, method, *inputs, **kwargs): if method == '__call__': _index = None scalars = [] for input in inputs: if isinstance(input, Number): scalars.append(input) elif isinstance(input, self.__class__): scalars.append(input.data) if _index is not None: if _index != input.index: raise TypeError("inconsistent sizes") else: _index = input.index else: return NotImplemented return self.__class__(self.name, ufunc(*scalars, **kwargs), _index) else: return NotImplemented def __array_function__(self, func, types, args, kwargs): if func not in HANDLED_FUNCTIONS: return NotImplemented if not all(issubclass(t, self.__class__) for t in types): return NotImplemented return HANDLED_FUNCTIONS[func](*args, **kwargs) @implements(np.mean) def mean(arr): return np.mean(arr.data) @implements(np.sum) def sum(arr): return np.sum(arr.data) @implements(np.where) def where(condition, x, y): if isinstance(x, myClass): x = x.data if isinstance(y, myClass): y = y.data return np.where(condition, x, y)
错误信息
调用np.where(myclassobj > 1, 0, 1)时触发如下错误:
File ~test.py:74, in where(condition, x, y) 72 if isinstance(y, myClass): 73 y = y.data ---> 74 return np.where(condition, x, y) File ~test.py:53, in myClass.__array_function__(self, func, types, args, kwargs) 51 if func not in HANDLED_FUNCTIONS: 52 return NotImplemented ---> 53 if not all(issubclass(t, self.__class__) for t in types): 54 return NotImplemented 55 return HANDLED_FUNCTIONS[func](*args, **kwargs) RecursionError: maximum recursion depth exceeded
有没有人有实现自定义数组np.where方法的经验?我参考了dask和cupy的实现但难以理解。
环境:Python 3.10,NumPy 1.26.4
内容的提问来源于stack exchange,提问作者john-projs
相关产品推荐
相关产品推荐

