如何将自定义myTensor类对象传入numpy的dot函数实现目标运算效果
实现方案
Numpy 的通用函数支持自定义类通过实现 __array_ufunc__ 特殊方法适配调用逻辑,你只需要在 myTensor 类中补充该方法,专门处理 np.dot 的调用场景即可。
修改后的完整类代码如下:
import numpy as np class myTensor: def __init__(self,data): self.data=np.array(data) self.parent=[] def __array_ufunc__(self, ufunc, method, *inputs, **kwargs): # 匹配np.dot调用场景 if ufunc == np.dot and method == '__call__' and len(inputs) == 2: a, b = inputs if isinstance(a, myTensor) and isinstance(b, myTensor): # 计算点积生成新实例 dot_res = np.dot(a.data, b.data) out = myTensor(dot_res) out.parent = [a, b] return out # 其他运算场景走默认逻辑,可按需扩展 return NotImplemented
效果验证
运行你给出的示例代码即可得到预期结果:
t1=myTensor([1,2]) t2=myTensor([3,4]) t3=np.dot(t1,t2) print(t3.data) # 输出 11 print(t3.parent == [t1, t2]) # 输出 True
内容的提问来源于stack exchange,提问作者AZ96
相关产品推荐
相关产品推荐

