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

如何为元素是自定义类实例的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 05:03:20