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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 22:46:12