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

重写Numpy的__array_function__方法时的性能损耗问题

问题

我编写了一个类Vector,它继承自np.ndarray,行为和原生数组类似,还额外添加了几何引擎需要的属性和方法(此处省略)。为了让调用np.dot时返回Vector对象,我重写了__array_function__方法,但做基准测试时发现性能比原生np.array差很多:

最小可复现代码

import numpy as np
from timeit import timeit


class Vector(np.ndarray):

    def __new__(cls, input_array):
        return np.array(input_array).view(cls)

    def __array_function__(self, func, types, args, kwargs):
        if func == np.dot:
            out = np.dot(np.asarray(args[0]), np.asarray(args[1]))
            return out.view(Vector)

基准测试代码

v = Vector([1, 1, 1])
I = np.identity(3)

print(type(np.dot(I, v)))  # 确保返回正确类型

# 创建np.array和Vector对象
A = np.random.random((100, 3))
V = A.view(Vector)

# 对比np.dot速度
print(timeit(lambda: np.dot(I, A.T)))
print(timeit(lambda: np.dot(I, V.T)))

输出结果

<class '__main__.Vector'>
1.207045791001292
2.063941927997803

性能损耗高达70%,请问这是正常现象吗?我有没有操作错误?有没有针对np.dot和np.cross的解决办法?如果没有的话,我可能得放弃这个自定义类。


分析与解决方法

1. 性能损耗是否正常?

这种程度的性能损耗是正常但可优化的。核心原因是__array_function__本身会带来numpy函数分发的额外开销,再加上你之前的代码里手动转换参数、重复调用np.dot后再转类型,多了几层不必要的操作,进一步放大了性能差距。

2. 操作是否有误?

核心逻辑没问题,但实现方式不够高效:

  • 手动调用np.asarray转换参数是多余的,Vector本身就是np.ndarray子类,直接传递给np.dot就能正常处理;
  • 没有利用numpy原生的__array_function__调度逻辑,而是手动重新实现np.dot调用,额外增加了函数调用开销。

3. 针对np.dot和np.cross的优化方案

方案一:优化__array_function__实现

修改__array_function__,让numpy先处理原生计算逻辑,最后只做一次类型转换,避免手动转换参数的开销:

class Vector(np.ndarray):

    def __new__(cls, input_array):
        return np.array(input_array).view(cls)

    def __array_function__(self, func, types, args, kwargs):
        # 只处理关注的两个函数
        if func in (np.dot, np.cross):
            # 调用父类逻辑,让numpy用原生高效路径计算
            result = super().__array_function__(func, types, args, kwargs)
            # 将结果转为Vector返回(标量结果直接返回,不转类型)
            return result.view(Vector) if np.ndim(result) > 0 else result
        # 其他函数交给父类默认处理
        return super().__array_function__(func, types, args, kwargs)

这种写法能大幅降低性能损耗,因为numpy会直接用原生的优化逻辑执行计算,最后只做一次轻量的类型转换。

方案二:直接重载实例方法(性能最优)

如果只关注dot和cross,可以直接给Vector添加实例方法,绕过__array_function__的调度开销:

class Vector(np.ndarray):

    def __new__(cls, input_array):
        return np.array(input_array).view(cls)

    def dot(self, other):
        result = np.dot(self, np.asarray(other))
        return result.view(Vector) if np.ndim(result) > 0 else result

    def cross(self, other):
        result = np.cross(self, np.asarray(other))
        return result.view(Vector)

使用时直接调用v.dot(other)而非np.dot(v, other),性能几乎和原生数组一致。如果需要兼容np.dot的调用方式,可以同时保留优化后的__array_function__。

方案三:使用__array_ufunc__(针对底层操作)

__array_ufunc__是numpy处理通用函数的协议,开销略低于__array_function__,如果需要处理更多通用函数操作可以尝试:

class Vector(np.ndarray):

    def __new__(cls, input_array):
        return np.array(input_array).view(cls)

    def __array_ufunc__(self, ufunc, method, *inputs, **kwargs):
        if method == '__call__' and ufunc in (np.dot, np.cross):
            inputs = tuple(np.asarray(inp) for inp in inputs)
            result = ufunc(*inputs, **kwargs)
            return result.view(Vector) if np.ndim(result) > 0 else result
        return super().__array_ufunc__(ufunc, method, *inputs, **kwargs)

不过这个方案的收益不如前两个明显,适合需要扩展更多ufunc操作的场景。


内容的提问来源于stack exchange,提问作者CyrielN

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 10:45:18