如何为元素是自定义类实例的NumPy数组实现矩阵乘法
自定义类实例组成的NumPy数组乘法实现方案
问题描述
是否存在简便方法,可重写元素为已实现__mul__、__add__方法的自定义类实例的NumPy数组乘法运算?
示例代码如下:
import numpy class A: def __init__(self): ... def __add__(self): ... def __mul__(self): ... def __repr__(self): ... instance = numpy.array([[A(...), A(...)],[A(...),A(...)]]) instance @ instance
运行上述代码会抛出如下错误:
TypeError: ufunc 'matmul' not supported for the input types, and the inputs could not be safely coerced to any supported types according to the casting rule ''safe''
已尝试通过实现__array_ufunc__、__matmul__方法解决该报错,但未找到可参考的用法示例。
需求优先级:
- 优先希望复用NumPy数组的原生方法提升使用便捷性
- 如果无法复用原生方法,仅实现
__matmul__方法支持矩阵乘法也可满足要求
可行实现方案
方案1:实现__array_ufunc__对接NumPy原生运算机制
NumPy的@运算符对应numpy.matmul ufunc,只要自定义类实现__array_ufunc__方法,就能被NumPy的运算分发逻辑识别,不需要修改数组侧逻辑,所有存储该类实例的object数组都能直接使用原生运算符。
可参考以下实现:
import numpy as np class A: def __init__(self, val): self.val = val def __add__(self, other): if isinstance(other, A): return A(self.val + other.val) return NotImplemented def __mul__(self, other): if isinstance(other, A): return A(self.val * other.val) return NotImplemented def __repr__(self): return f"A({self.val})" def __array_ufunc__(self, ufunc, method, *inputs, **kwargs): if method != "__call__": return NotImplemented # 统一将输入中的NumPy数组转为Python原生列表,方便逐元素调用类的运算方法 def parse_input(x): if isinstance(x, np.ndarray): return x.tolist() return x input_list = [parse_input(item) for item in inputs] # 按不同ufunc分发运算逻辑 if ufunc is np.matmul: a, b = input_list n_row = len(a) n_col = len(b[0]) n_mid = len(b) res = [] for i in range(n_row): row = [] for j in range(n_col): # 矩阵乘:对应位置相乘累加 elem = a[i][0] * b[0][j] for k in range(1, n_mid): elem = elem + a[i][k] * b[k][j] row.append(elem) res.append(row) return np.array(res, dtype=object) # 按需扩展支持逐元素乘、逐元素加等其他运算 if ufunc is np.multiply: a, b = input_list return np.array([[ai*bi for ai, bi in zip(row_a, row_b)] for row_a, row_b in zip(a,b)], dtype=object) if ufunc is np.add: a, b = input_list return np.array([[ai+bi for ai, bi in zip(row_a, row_b)] for row_a, row_b in zip(a,b)], dtype=object) return NotImplemented
测试代码:
# 注意必须显式指定dtype=object,否则NumPy会尝试将实例转为内置数值类型 instance = np.array([[A(1), A(2)],[A(3),A(4)]], dtype=object) print(instance @ instance) # 输出:[[A(7), A(10)], [A(15), A(22)]],和普通数值矩阵乘结果一致
注意事项:
- 定义存储自定义类实例的数组时,必须显式指定
dtype=object,否则NumPy会默认尝试将实例转换为内置数值类型,提前触发类型转换错误 __array_ufunc__方法内可按需扩展支持更多NumPy运算逻辑,无需修改数组侧实现,复用性更强
方案2:自定义NumPy数组子类重写__matmul__
如果不需要兼容全量NumPy ufunc,仅需要支持@矩阵乘法,可以继承np.ndarray实现自定义数组子类,重写__matmul__方法即可,逻辑更简单:
class CustomObjArray(np.ndarray): def __matmul__(self, other): a = self.tolist() b = other.tolist() n_row = len(a) n_col = len(b[0]) n_mid = len(b) res = [] for i in range(n_row): row = [] for j in range(n_col): elem = a[i][0] * b[0][j] for k in range(1, n_mid): elem = elem + a[i][k] * b[k][j] row.append(elem) res.append(row) return np.array(res, dtype=object).view(CustomObjArray) # 测试 instance = np.array([[A(1), A(2)],[A(3),A(4)]], dtype=object).view(CustomObjArray) print(instance @ instance) # 输出和方案1一致
内容的提问来源于stack exchange,提问作者ChoMedit
相关产品推荐
相关产品推荐

