如何让numpy.concatenate返回自定义numpy数组子类而非普通ndarray
如何让numpy.concatenate返回自定义numpy数组子类而非普通ndarray
嘿,这个问题我之前也碰到过!numpy的数组子类确实需要额外处理一些数组操作的返回类型,光靠__new__是不够的,你猜的没错,__array_finalize__是关键之一,但我们还需要结合__array_function__来拦截concatenate这类函数的调用,确保返回的是我们的子类实例。
先给你直接上修改后的可运行代码:
import numpy as np class BreakfastArray(np.ndarray): def __new__(cls, n=1): dtypes=[("waffles", int), ("eggs", int)] obj = np.zeros(n, dtype=dtypes).view(cls) return obj def __array_finalize__(self, obj): # 这个方法是numpy子类的初始化钩子,用于同步父实例的状态 # 目前我们的类没有自定义属性,所以保持空实现也能保证类型继承的基础逻辑 if obj is None: return def __array_function__(self, func, types, args, kwargs): # 拦截numpy.concatenate函数的调用 if func is np.concatenate: # 先调用原生concatenate得到普通ndarray结果 result = func(*args, **kwargs) # 将结果转换为我们的BreakfastArray子类 return result.view(type(self)) # 其他numpy函数交给父类的默认逻辑处理 return super().__array_function__(func, types, args, kwargs) # 测试验证 b1 = BreakfastArray(n=1) b2 = BreakfastArray(n=2) con_b1b2 = np.concatenate([b1, b2]) print(b1.__class__, con_b1b2.__class__)
运行这段代码,输出就会是<class '__main__.BreakfastArray'> <class '__main__.BreakfastArray'>,完全符合你的需求!
简单解释下关键部分:
__array_finalize__:当通过视图、切片或数组操作创建新实例时,numpy会自动调用这个方法。它是子类继承状态的核心,后续如果给BreakfastArray添加自定义属性,就需要在这里把父实例的属性复制到新实例中,避免属性丢失。__array_function__:这是numpy的函数调度协议,用来让子类接管特定numpy函数的执行逻辑。我们在这里拦截concatenate,先拿到原生结果再转成子类,实现了自动返回自定义类型的效果。
如果之后你的子类有了自定义属性,一定要记得在__array_finalize__里同步这些属性哦!
备注:内容来源于stack exchange,提问作者I.P. Freeley
相关产品推荐
相关产品推荐

