如何在Numpy中支持dtype=object对象数组的矩阵乘法?
Got it, let's tackle this problem. So you've got NumPy arrays filled with your custom Ciphertext instances (with __add__ and __mul__ overloaded), and while element-wise ops and broadcasting work great, matrix multiplication (like @ or np.dot) isn't behaving as expected. That makes sense—NumPy's built-in matrix multiplication routines are optimized for numeric types, and they don't automatically map to your custom object's operators. Here's how to fix it:
1. 理解核心问题
NumPy的默认矩阵乘法逻辑是针对数值类型编写的,当遇到dtype=object的数组时,它无法识别你重载的__mul__和__add__来完成矩阵乘法所需的"乘加"操作。所以我们需要手动实现这个逻辑,或者引导NumPy使用我们自定义的运算符。
2. 手动实现矩阵乘法(基础版)
最直接的方式是写一个函数,遍历数组的行和列,用你自定义的运算符完成点乘求和。这里假设你的Ciphertext类支持加法单位元(比如可以初始化一个代表0的实例):
首先,补全示例代码上下文方便测试:
import numpy as np # 假设的Ciphertext类(包含你提到的运算符重载) class Ciphertext: def __init__(self, value): self.value = value def __add__(self, other): return Ciphertext(self.value + other.value) def __mul__(self, other): return Ciphertext(self.value * other.value) def __repr__(self): return f"Ciphertext({self.value})" # 假设的加密构建器 class EncryptionBuilder: def encrypt_as_array(self, arr): return np.array([Ciphertext(x) for x in arr.flat], dtype=object).reshape(arr.shape) # 初始化测试数据 builder = EncryptionBuilder() encrypted_a = builder.encrypt_as_array(np.array([[1, 2], [3, 4]])) encrypted_b = builder.encrypt_as_array(np.array([[5, 6], [7, 8]]))
然后实现基础版矩阵乘法:
def object_matmul(a, b): # 仅处理二维数组(可根据需求扩展到更高维) if a.ndim != 2 or b.ndim != 2: raise ValueError("This function currently supports 2D arrays only") if a.shape[1] != b.shape[0]: raise ValueError(f"Shape mismatch: {a.shape} cannot multiply with {b.shape}") # 初始化结果数组 result_shape = (a.shape[0], b.shape[1]) result = np.empty(result_shape, dtype=object) # 遍历计算每个元素:行向量点乘列向量 for i in range(result_shape[0]): for j in range(result_shape[1]): # 初始化点积为0的Ciphertext实例 dot_product = Ciphertext(0) for k in range(a.shape[1]): dot_product += a[i, k] * b[k, j] result[i, j] = dot_product return result # 测试调用 matmul_result = object_matmul(encrypted_a, encrypted_b) print(matmul_result) # 输出: # [[Ciphertext(19) Ciphertext(22)] # [Ciphertext(43) Ciphertext(50)]]
3. 优化版本(利用Numpy广播加速)
三重循环对于大数组来说效率较低,我们可以利用Numpy的广播机制减少循环次数,同时借助np.sum完成求和。不过这里需要给Ciphertext类添加__radd__重载,因为np.sum默认会从数值0开始累加,需要支持0 + Ciphertext的操作:
首先更新Ciphertext类:
class Ciphertext: # 保留原有方法... def __radd__(self, other): # 处理sum的初始值(数值0)的情况 if other == 0: return self return self + other
然后实现优化版矩阵乘法:
def optimized_object_matmul(a, b): if a.ndim != 2 or b.ndim != 2: raise ValueError("Supports 2D arrays only") if a.shape[1] != b.shape[0]: raise ValueError(f"Shape mismatch: {a.shape} cannot multiply with {b.shape}") # 广播得到元素级乘积的三维数组:(a行数, 公共维度, b列数) element_products = a[:, :, np.newaxis] * b[np.newaxis, :, :] # 沿着公共维度求和,得到最终矩阵 result = np.sum(element_products, axis=1) return result # 测试调用 optimized_result = optimized_object_matmul(encrypted_a, encrypted_b) print(optimized_result) # 输出和基础版一致,但速度快很多
4. 支持@运算符(进阶)
如果想直接用@运算符(就像普通数值数组一样),可以创建一个继承自np.ndarray的自定义数组类,重载__matmul__方法:
class EncryptedArray(np.ndarray): def __new__(cls, input_array): # 从输入的object数组创建EncryptedArray实例 obj = np.asarray(input_array, dtype=object).view(cls) return obj def __matmul__(self, other): # 调用优化后的矩阵乘法函数 return optimized_object_matmul(self, other) # 使用方式 encrypted_arr_a = EncryptedArray(builder.encrypt_as_array(np.array([[1,2],[3,4]]))) encrypted_arr_b = EncryptedArray(builder.encrypt_as_array(np.array([[5,6],[7,8]]))) # 直接用@运算符 final_result = encrypted_arr_a @ encrypted_arr_b print(final_result)
总结
- NumPy原生矩阵乘法不支持
dtype=object数组的自定义运算符,因为它底层是针对数值类型的优化实现。 - 基础版手动循环适合小数组或者快速验证,优化版借助广播和
np.sum能提升大数组的运算效率。 - 自定义数组类可以让你像使用普通NumPy数组一样用
@运算符,代码更简洁。
内容的提问来源于stack exchange,提问作者Overholt

