如何实现NumPy数组与自定义Vector3d迭代器的乘法运算?
自定义Vector3d类与NumPy数组乘法报错的解决方法
问题重现
自定义的Vector3d类代码如下:
from collections.abc import Iterable class Vector3d(object): def __init__(self, *args): if len(args)==3: self.x = args[0] self.y = args[1] self.z = args[2] else: self.x, self.y, self.z = args[0] def _apply(self, func, o): if isinstance(o, Iterable): xyz = [self.x, self.y, self.z] for i,j in enumerate(o): xyz[i] = func(xyz[i], j) return Vector3d(xyz) else: return Vector3d(func(self.x, o), func(self.y, o), func(self.z, o)) def __getitem__(self, i): if i=='x' or i==0: return self.x if i=='y' or i==1: return self.y if i=='z' or i==2: return self.z def __repr__(self): return "Vector3d({}, {}, {})".format(self.x, self.y, self.z) def __iter__(self): return iter([self.x,self.y,self.z]) def __neg__(self): return Vector3d(-self.x, -self.y, -self.z) def __add__(self, o): return self._apply(lambda x,y: x+y, o) def __sub__(self, o): return self._apply(lambda x,y: x-y, o) def __mul__(self, o): return self._apply(lambda x,y: x*y, o) def __truediv__(self, o): return self._apply(lambda x,y: x/y, o) def __eq__(self, o): return all(self._apply(lambda x,y: x==y, o)) def __ne__(self, o): return any(self._apply(lambda x,y: x!=y, o))
执行与NumPy数组的乘法时触发错误:
>>> import numpy as np >>> a = np.array([[1,2,3],[4,5,6],[5,6,7]]) >>> a*Vector3d(1,2,3) Traceback (most recent call last): File "<stdin>", line 1, in <module> TypeError: unsupported operand type(s) for *: 'int' and 'Vector3d'
问题原因
NumPy数组采用逐元素运算逻辑:执行a * v时,NumPy会尝试将v与数组a的每个元素做乘法。现有Vector3d类仅实现了__mul__(处理v * o的场景),未实现__rmul__(处理o * v的场景,这里o是NumPy数组的int类型元素),导致NumPy找不到int与Vector3d的乘法实现,从而报错。
另外原_apply方法用isinstance(o, Iterable)判断会把NumPy数组误判为可迭代对象,走错误分支进一步触发问题。
解决方法
方法1:实现__rmul__适配NumPy运算
在Vector3d类中添加__rmul__方法,同时处理标量和NumPy数组场景:
import numpy as np from collections.abc import Iterable class Vector3d(object): # 保留原有所有方法... def __rmul__(self, o): # 处理标量与Vector3d的乘法(对应NumPy数组元素的运算) if np.isscalar(o): return Vector3d(self.x * o, self.y * o, self.z * o) # 处理NumPy数组的逐元素广播乘法 elif isinstance(o, np.ndarray): vec_np = np.array([self.x, self.y, self.z]) return o * vec_np[np.newaxis, :] if o.ndim == 2 else o * vec_np # 其他场景复用_apply逻辑 else: return self._apply(lambda x,y: y*x, o)
执行a * Vector3d(1,2,3)会返回预期的NumPy数组:
array([[ 1, 4, 9], [ 4, 10, 18], [ 5, 12, 21]])
方法2:实现__array__让Vector3d可转为NumPy数组
添加__array__方法,让NumPy自动识别并转换Vector3d为数组:
class Vector3d(object): # 保留原有所有方法... def __array__(self): return np.array([self.x, self.y, self.z])
之后可以直接转换后运算:
>>> a * np.array(Vector3d(1,2,3)) array([[ 1, 4, 9], [ 4, 10, 18], [ 5, 12, 21]])
方法3:优化_apply的类型判断逻辑
修改_apply方法,优先判断NumPy数组,避免误判可迭代对象:
def _apply(self, func, o): # 优先处理NumPy数组 if isinstance(o, np.ndarray): vec_np = np.array([self.x, self.y, self.z]) result = func(vec_np, o) return result if isinstance(result, np.ndarray) else Vector3d(result) # 排除字符串/字节类型的可迭代对象,避免误判 elif isinstance(o, Iterable) and not isinstance(o, (str, bytes)): xyz = [self.x, self.y, self.z] for i,j in enumerate(o): xyz[i] = func(xyz[i], j) return Vector3d(xyz) else: return Vector3d(func(self.x, o), func(self.y, o), func(self.z, o))
配合__rmul__方法,可实现双向乘法运算的兼容性。
内容的提问来源于stack exchange,提问作者IsaacMak
相关产品推荐
相关产品推荐

