重写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

