ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

Python 3.12魔法方法__matmul__:运算符重载与矩阵乘法实现

2026/9/28 14:23:24 拓冰建站 浏览量
Python 3.12魔法方法__matmul__:运算符重载与矩阵乘法实现 1. Python 3.12魔法方法系列深入__matmul__运算符重载这期是这个系列的第47篇把__matmul__这个魔法方法单独拎出来聊。你要是写过矩阵运算相关的代码大概率见过这个运算符。在 Python 3.5 引入作为矩阵乘法运算符之后__matmul__就成了实现这一行为的关键入口。到了 Python 3.12虽然解释器内部做了不少字节码层面的优化但这个魔法方法本身的使用方式和设计思路依然稳定而且在实际工程中发挥的作用越来越大。1.1 从矩阵乘法说起为什么需要专门一个运算符很多从其他语言转过来的开发者第一次看到a b都会愣一下心想这不就是装饰器的语法吗怎么能用来做乘法这事要从数值计算的需求说起。在科学计算、机器学习、深度学习领域矩阵乘法是最基础也是最频繁的操作之一。如果只能用*去做矩阵乘法就得区分“逐元素相乘”和“矩阵内积”这两种完全不同的语义代码可读性会很差。Python 社区在 PEP 465 中正式引入了运算符专门用于矩阵乘法。__matmul__就是 Python 解释器在遇到a b时自动调用的魔法方法它让我们自定义的类能够支持这个运算符实现真正的“语法级”支持。到了 Python 3.12这个机制依旧稳定没有出现破坏性变更这本身就说明当初的设计足够前瞻。1.2 这个魔法方法到底做了什么三个相关方法的完整配合要理解__matmul__不能只看它本身。Python 的二元运算符有一套完整的协议也不例外。除了__matmul__还有两个“兄弟方法”需要一起了解__matmul__(self, value)处理self value也就是左操作数发起运算的情况__rmatmul__(self, value)处理value self也就是右操作数发起但左操作数不支持矩阵乘法的场景__imatmul__(self, value)处理self value也就是原地矩阵乘法的场景我见过不少初学者只实现了__matmul__然后在写2 obj这样的表达式时收到TypeError: unsupported operand type(s)这就是因为没实现__rmatmul__。形成一个完整的运算符支持三个方法通常要一起考虑。方法名触发的表达式说明__matmul__a b常规矩阵乘法__rmatmul__b a且b不支持时反向矩阵乘法__imatmul__a b原地矩阵乘法2. 手写实现一个矩阵类完整解决__matmul__的设计难题光讲概念没有说服力直接写代码。这里我实现一个支持运算的轻量级矩阵类Mat。目标有两个一是把__matmul__和它的兄弟方法完整实现一遍二是通过这个过程把运算符重载的设计逻辑讲透。2.1 基础框架数据存储与维度校验class Mat: 一个支持矩阵乘法运算符的轻量级矩阵类 def __init__(self, data): # 使用列表的列表来存储矩阵数据 self._data [row[:] for row in data] # 深拷贝防止外部修改 self._validate() def _validate(self): # 校验数据格式需要是二维列表且每一行长度一致 if not self._data or not all(isinstance(row, list) for row in self._data): raise TypeError(必须传入非空二维列表) row_len len(self._data[0]) if any(len(row) ! row_len for row in self._data): raise ValueError(每一行的列数必须一致) property def shape(self): return (len(self._data), len(self._data[0]) if self._data else 0) def __repr__(self): return fMat({self._data})这里做了一个关键决定在构造时进行深拷贝。为什么矩阵类通常应该具备值语义也就是说赋值和传参不应该像列表那样直接共享底层数据。如果你的矩阵类可以被外部意外修改内部存储用运算时出现问题会非常难排查。虽然深拷贝会损失一些性能但换来的安全性非常值得。2.2 核心实现matmul的完整逻辑接下来是实现矩阵乘法的核心方法def __matmul__(self, other): if isinstance(other, Mat): return self._matmul_mat(other) elif isinstance(other, (int, float)): # 支持矩阵和标量的乘法 return Mat([[val * other for val in row] for row in self._data]) else: return NotImplemented注意这里有几个关键设计。首先当遇到不支持的类型时返回的是NotImplemented而不是直接抛异常。这是 Python 运算符重载的重要规则如果你直接抛TypeError那么当左操作数不支持而右操作数支持时解释器就没有机会再去尝试右操作数的反向方法了。返回NotImplemented是在告诉 Python我不处理这种情况请你试试别的方案。_matmul_mat内部实现如下def _matmul_mat(self, other): # 维度检查左边矩阵的列数必须等于右边矩阵的行数 if self.shape[1] ! other.shape[0]: raise ValueError(f矩阵维度不匹配: {self.shape} {other.shape}) # 常规的三重循环矩阵乘法 result [] for i in range(self.shape[0]): # 左矩阵的行 row [] for j in range(other.shape[1]): # 右矩阵的列 total 0 for k in range(self.shape[1]): # 左矩阵的列/右矩阵的行 total self._data[i][k] * other._data[k][j] row.append(total) result.append(row) return Mat(result)维度检查放在运算之前非常重要这不仅是为了保证语义正确更是为了给出清晰的报错信息。Python 的 numpy 库在维度不匹配时会抛出类似ValueError: matmul: Input operand 1 has a mismatch in its core dimension 0的信息我们在自定义类里也应该提供同等质量的错误提示。2.3 补齐反向方法和原地方法完整运算符协议的必修课只有__matmul__是不够的。完整的运算符支持还需要__rmatmul__和__imatmul__def __rmatmul__(self, other): # 当左操作数不是 Mat右操作数是 Mat 时触发 if isinstance(other, (int, float)): # 标量与矩阵的乘法利用交换律 return Mat([[val * other for val in row] for row in self._data]) return NotImplemented def __imatmul__(self, other): # 原地矩阵乘法直接修改自身数据 result self other # 复用 __matmul__ 的逻辑 self._data result._data return self__rmatmul__的触发场景值得多说一句。当你写2 mat时Python 首先检查整数2有没有__matmul__方法。整数没有这个方法于是 Python 检查右边mat有没有__rmatmul__。如果在右边找到了就调用它。这个机制可以避免很多类型判断的麻烦也是 Python 运算符协议相对完善的体现。__imatmul__的原地操作语义上要注意它和非原地操作最大的区别在于是否返回新对象。a a b会创建一个新的矩阵对象而a b最好直接修改a本身。这么做不只是性能上的考虑更关系到代码的可预测性——如果用户以为修改了原对象结果返回了一个新对象并重新绑定给变量行为上可能没问题但如果你用一个变量引用了原矩阵它没被修改这时行为就有差异了。3. 矩阵乘法背后的数学原理与Python 3.12的字节码优化到了 Python 3.12你写a b时解释器内部发生了一些有意思的变化。这一节把数学层面的运算逻辑和解释器层面的调用机制都说清楚。3.1 矩阵乘法的数学定义为什么计算顺序非交换先复习一个基础概念。设矩阵 A 的形状是(m, n)矩阵 B 的形状是(n, p)那么 C A B 的形状是(m, p)其中 C[i][j] 的计算公式是C[i][j] sum(A[i][k] * B[k][j], k 0 to n-1)注意这里有个前提A 的列数必须等于 B 的行数都是 n。这也是上一节维度检查的数学基础。矩阵乘法不满足交换律A B 和 B A 大多数情况下结果完全不同甚至维度都对不上。这也是为什么运算符和*运算符需要被严格区分——*在矩阵场景下的交换律是成立的而往往不成立。知道了计算逻辑再看刚才的_matmul_mat三重循环就再清晰不过了。最内层的for k in range(self.shape[1])就是在做那个累加和。3.2 Python 3.12 对 BINARY_OP 字节码的优化Python 3.12 在执行a b时实际上会编译成BINARY_OP字节码并在该字节码中传入一个操作符参数Python 3.11 以前是BINARY_MATRIX_MULTIPLY这种独立字节码。这个改变的核心意义在于简化了 Python 虚拟机的指令集但对外层行为没有影响。对于使用__matmul__的开发者来说意味着你的代码在 Python 3.12 下基本可以直接运行不需要修改。但有一点值得注意Python 3.12 对二进制操作在字节码层面的缓存机制做了改进——CPython 解释器会缓存类型检查的结果这就意味着在同一个循环里反复执行同一类型的运算可以跳过重复的方法查找过程。如果你的代码里有一个大循环做矩阵乘法相比之前版本会有更好的运行时性能。3.3 代码demo验证Python 3.12下的运算行为写段代码验证一下在 Python 3.12 下用运算符的实际行为。假设还是用刚才的Mat类# 在 Python 3.12 环境运行 import sys print(sys.version) # 确认是 3.12 系列 m1 Mat([[1, 2], [3, 4]]) m2 Mat([[5, 6], [7, 8]]) # 使用 运算符 result m1 m2 print(result) # 输出: Mat([[19, 22], [43, 50]]) # 验证手动计算 # [1*5 2*7, 1*6 2*8] [19, 22] # [3*5 4*7, 3*6 4*8] [43, 50] # 反向运算 result2 2 m1 print(result2) # 输出: Mat([[2, 4], [6, 8]]) # 原地运算 m3 Mat([[1, 2], [3, 4]]) m3 m2 print(m3) # 输出: Mat([[19, 22], [43, 50]]) # 维度不匹配的情况 m4 Mat([[1, 2, 3], [4, 5, 6]]) # 形状 (2, 3) try: m5 m1 m4 # (2,2) (2,3) 形状不匹配 except ValueError as e: print(f捕获错误: {e}) # 输出: 捕获错误: 矩阵维度不匹配: (2, 2) (2, 3)这个验证过程很直观能看出运算符的行为完全符合 PEP 465 的设计预期。你可以把自己的自定义类中的__matmul__和 numpy 的 ndarray 结合在一起实现更复杂的运算逻辑。4. 实战优化性能对比、运算符链与复杂场景处理如果只是把__matmul__实现出来就能满足基本需求了。但实际工程里性能、可读性和扩展性往往更重要。这一节分享一些真实的优化思路和复杂场景的应对策略。4.1 用列表推导优化一个替代三重循环的写法前面用三重循环实现了_matmul_mat目的纯粹是为了把数学过程展示清楚。实际写代码时你不必写得这么原始。Python 的列表推导不仅代码量更少而且由于内部循环在 C 层面优化通常会有更好的性能表现。def _matmul_mat_fast(self, other): if self.shape[1] ! other.shape[0]: raise ValueError(f矩阵维度不匹配: {self.shape} {other.shape}) # 先转置右矩阵让内层循环变成连续内存访问 other_T list(zip(*other._data)) return Mat([ [sum(a * b for a, b in zip(row, col)) for col in other_T] for row in self._data ])转置右矩阵是一个常见的优化手段。原始的_matmul_mat中访问other._data[k][j]时每次k变化都跳到不同的行内存访问不连续CPU 缓存命中率低。转置后col是一个连续的行列表计算zip(row, col)时两个序列都是顺序访问的对 Python 底层实现的友好程度更高。当然如果真要看性能就不要用 Python 原生做了。numpy 的底层是高度优化的 BLAS 库性能差距是数量级的。但在自定义类里了解这些优化思想远比直接调用一个高效工具更有价值。4.2 运算符链理解Python的调用顺序Python 处理a b c时运算顺序是(a b) c。也就是说先计算a b的结果再和c相乘。这个逻辑对习惯数学运算的人很自然但在实现时要注意一个问题中间结果的数据类型会影响后续运算。假如a b返回的是Mat那继续 c没有问题。但是如果a b返回的恰好是一个列表比如你在__matmul__里返回了一个嵌套列表而不是新的Mat对象那么result c就会因为列表没有__matmul__方法而报错。所以__matmul__方法需要保持一致的数据类型返回策略这是运算符链能够优雅工作的关键。# 验证运算符链 m1 Mat([[1, 0], [0, 1]]) # 单位矩阵 m2 Mat([[2, 3], [4, 5]]) m3 Mat([[1, 1], [1, 1]]) # (m1 m2) m3 result m1 m2 m3 print(result) # 输出: Mat([[5, 5], [9, 9]])4.3 与numpy共存你的类如何与ndarray协同大部分做数值计算的 Python 项目里numpy 几乎不可避免。自定义的Mat类如果完全不考虑和 numpy 协同局限性会很大。一个务实的方案是让你的__matmul__方法支持 numpy 数组作为other参数。def __matmul__(self, other): if isinstance(other, Mat): return self._matmul_mat(other) elif hasattr(other, __array__): # 兼容 numpy ndarray 以及其他数组类型 import numpy as np arr np.array(self._data) np.asarray(other) return Mat(arr.tolist()) elif isinstance(other, (int, float)): return Mat([[val * other for val in row] for row in self._data]) else: return NotImplemented这么做的好处在于你的自定义类能和 numpy 数组自由混用提升了代码灵活性。需要注意import numpy as np最好放在方法内部或模块顶部整体导入避免因为 import 本身有较大开销导致频繁调用时性能受损。5. 实战小专题在项目里应用matmul的常见坑与最佳实践这一节整理我实际踩过的坑和总结出来的最佳实践直接给结论和解决方案。5.1 常见坑不返回NotImplemented导致的类型错误这是最典型的坑。假设你在__matmul__里没有返回NotImplemented而是直接写def __matmul__(self, other): if not isinstance(other, Mat): raise TypeError(不支持的类型) # ...那么当你执行some_object mat而some_object恰好也实现了__rmatmul__时Python 就永远不会走到__rmatmul__的调用。这会让你的反向扩展失效。正确的做法是def __matmul__(self, other): if isinstance(other, Mat): return self._matmul_mat(other) elif isinstance(other, (int, float)): return Mat([[val * other for val in row] for row in self._data]) # 其它情况返回 NotImplemented return NotImplemented再补充一点如果两个对象都不支持某项运算Python 最终会抛出TypeError: unsupported operand type(s) for : Mat and XXX。这个报错信息来自解释器层面而不是你的自定义异常所以不用额外处理。5.2 只读与可变性的权衡浅拷贝、深拷贝、还是视图矩阵类要不要可变numpy 默认是可变且可以进行原地操作的你可以在同一个数组上反复做运算这对内存敏感的场景非常友好。但自定义类里的可变性需要谨慎设计。我建议遵循一个原则通过__init__构建的数据要与其他实例隔离但实例内部可以有一个_data的可变引用。那样的话可以原地修改内部数据而不产生新对象但外部无法通过传入的列表直接篡改数据。用深拷贝还是浅拷贝取决于你的数据规模——浅拷贝能降低构造开销适合大型矩阵但外部修改原始列表会波及矩阵对象内部。没有统一的正确答案关键是想清楚并保持一致性。5.3 性能讨论纯Python矩阵乘法的边界在哪里用纯 Python 写矩阵乘法性能上限其实是可预估的。一个(100, 100)的矩阵乘法需要做 100 万次内层乘加操作纯 Python 大约需要 0.1 秒左右如果用自定义的Mat类再带一层 Python 方法调用可能要 0.3 秒以上。而 numpy 在同样规模下几乎能到毫秒级甚至更快。所以我的建议是如果你的矩阵规模超过 100x100或者运算在一个大循环里频繁发生不要自己实现矩阵乘法老老实实调用 numpy。自定义类里实现__matmul__的价值在于在特定业务逻辑中定义“矩阵乘法”语义的另一种含义让自定义对象能无缝参与表达式提高可读性作为加深理解 Python 数据模型的练习理解了“什么时候该自己实现”和“什么时候该调库”本身就是工程能力的体现。6. 常见问题与排查技巧实录最后整理一份快速排查指南把实际操作中遇到的问题整理成对照表。问题现象可能原因解决方案报错TypeError: unsupported operand type(s) for 左操作数的__matmul__和右操作数的__rmatmul__都没有实现或都返回 NotImplemented分别实现两个方法确保返回 NotImplemented 而不是抛异常a b c报错中间结果的数据类型和预期不一致检查__matmul__是否始终返回同一类型对象矩阵维度不匹配但没有任何错误遗漏了维度检查导致计算结果错误在运算前明确检查self.shape[1] other.shape[0]a b没有按预期修改 a忘记实现__imatmul__Python 退化为a a b实现__imatmul__并在内部修改self._data使用混合 numpy 数组时报错__matmul__未兼容 numpy 数组类型添加对hasattr(other, __array__)的检查这里再分享一个排查技巧自定义类里的__matmul__实现有时默默返回了错误结果。遇到这种情况先做一个维度打印print(f运算维度: {self.shape} {other.shape})这个调试信息能快速定位问题出在维度判断还是计算逻辑上。我见过不少人在这一步省了力然后在结果完全不对时才回头排查花费的时间反而更多。还有一个细节如果你在实现里用了from __future__ import annotations要注意类型注解中引用自定义类时可能会遇到延迟求值的问题。如果你的类在模块内部定义或者涉及跨模块引用最好在类型注解中使用TypeVar或者Generic避免一些比较隐蔽的运行时错误。这不是__matmul__特有的问题但在类实现复杂时会碰到。最后__matmul__的文档字符串值得认真写。我见过最糟糕的一种实现是代码完全没有文档连维度限制都没写。一个稍微负责一点的实现至少应该写明运算规则self other。要求 self 的列数等于 other 的行数。 返回一个新 Mat 对象不会修改 self 和 other。 遇到不支持的 other 类型返回 NotImplemented。这种文档对团队协作非常重要——因为矩阵乘法本身的语义和习惯用法在数值计算领域里非常明确如果自己的实现对语义有偏差一定要写清楚否则合作者很容易被误导。