Python中实现支持NumPy/CuPy的规范化条件继承方案
问题:实现兼容NumPy/CuPy的动态继承数组类
我现有一个继承np.ndarray的类A:
class A(np.ndarray): def __new__(cls, foo: np.ndarray): # 对foo进行操作... x = np.array([1.]) + foo # 占位代码 return x def bar(self, x): # 其他操作...
我希望添加CPU/GPU无关的行为,通过检测输入是np.ndarray还是cupy.ndarray来动态继承其中一个类,大致思路如下:
class A: def __new__(cls, foo: Union[np.ndarray, cupy.ndarray]): xp = cupy.get_array_module(foo) # 根据foo返回`numpy`或`cupy` # 对foo进行操作... x = xp.array([1.]) + foo # 占位代码 return xp.__new__(cls, x) def bar(self, x): # 其他操作...
我尝试了以下实现方式,但觉得过于繁琐:
from abc import ABC import numpy as np import cupy as cp from typing import Union class BaseA(ABC): def bar(self, x): # 其他操作... pass class ANumpy(BaseA, np.ndarray): pass class ACupy(BaseA, cp.ndarray): pass class A: def __new__(cls, foo: Union[np.ndarray, cp.ndarray]): if isinstance(foo, cp.ndarray): return ACupy(foo) return ANumpy(foo)
请问如何优雅、规范地实现这一需求?
优雅实现方案
方案一:动态子类生成 + 方法复用
核心是在A的__new__中根据输入数组类型动态生成对应子类,自动绑定统一方法,避免手动维护多个子类:
import numpy as np import cupy as cp from typing import Union class A: def __new__(cls, foo: Union[np.ndarray, cp.ndarray]): xp = cp.get_array_module(foo) # 生成子类名称,区分NumPy/CuPy版本 subclass_name = f"A{'CuPy' if xp is cp else 'NumPy'}" # 缓存子类,避免重复创建 if subclass_name not in cls.__dict__: base_arr_cls = xp.ndarray # 动态创建子类,继承对应数组类并绑定自定义方法 subclass = type(subclass_name, (base_arr_cls,), { 'bar': cls.bar, # 可添加其他需要复用的方法 }) setattr(cls, subclass_name, subclass) # 预处理输入数组 x = xp.array([1.]) + foo # 返回动态子类的实例 return getattr(cls, subclass_name).__new__(getattr(cls, subclass_name), x) def bar(self, x): # 统一方法实现,自动适配CPU/GPU xp = cp.get_array_module(self) return xp.sum(self) + xp.sum(x)
优势
- 无需手动定义多个子类,代码精简
- 自定义方法只写一次,自动复用
- 子类首次创建后会缓存,不影响后续性能
方案二:混入类(Mixin)+ 工厂方法
用混入类封装统一业务逻辑,通过工厂方法动态选择继承的数组类,结构更清晰:
import numpy as np import cupy as cp from typing import Union class AMixin: def bar(self, x): # 统一行为实现 xp = cp.get_array_module(self) return xp.mean(self) * xp.mean(x) class A: @classmethod def create(cls, foo: Union[np.ndarray, cp.ndarray]): xp = cp.get_array_module(foo) # 动态创建继承自数组类和混入类的子类 subclass = type(f"A_{xp.__name__}", (xp.ndarray, AMixin), {}) # 预处理输入 x = xp.array([1.]) + foo return subclass.__new__(subclass, x)
优势
- 职责分离:混入类专注业务逻辑,工厂方法处理动态继承
- 同样避免了手动维护多个子类的繁琐
关键注意事项
- 所有数组操作必须通过
xp = cp.get_array_module(...)获取的模块执行,确保CPU/GPU无关性 - 动态生成的子类会自动继承原数组类的全部特性(索引、运算等),同时拥有自定义方法
- 如果需要处理数组视图/切片的状态传递,可以在动态子类中添加
__array_finalize__方法
内容的提问来源于stack exchange,提问作者Richie Bendall
相关产品推荐
相关产品推荐

