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

如何为自定义数组容器实现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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 12:34:56