提到 Python 魔术方法大家最熟悉的当然是__init__、__add__、__getitem__这些高频角色。今天聊的__matmul__相对冷门但它支撑的东西你绝对天天在用Python 3.5 之后新增的矩阵乘法运算符底层调用的就是它。你在 numpy 里写a b、在 torch 里写z x w背后的协议基础都是__matmul__。这篇文章我会把__matmul__从原理到实现完整拆一遍包括__rmatmul__和__imatmul__这两兄弟再带大家手写一个纯 Python 的矩阵类最后用斐波那契快速幂和最小二乘线性回归两个实战例子让你真正体会到这个魔法方法在科学计算里的分量。无论你是想搞懂 Python 运算符重载的初学者还是想给自己封装的类加上支持的老手这篇都能帮你少走不少弯路。1. 为什么 Python 要专门为矩阵乘法造一个运算符1.1的两重身份装饰器与矩阵乘法不少新手第一次看到会懵因为它在 Python 里已经有一种很常见的用途——装饰器语法比如staticmethod、classmethod、app.route。那到底算几个语法答案是同一符号两种完全独立的语法角色。装饰器出现在函数或类定义之前的行首本质上是把后面的函数作为参数传给装饰器函数而二元运算符出现在表达式中间语义是矩阵乘法。两者在语法层面不会混淆Python 解析器能清晰区分。也就是说不是替身它是真正的一等运算符。从 Python 3.5 开始a b这个表达式的标准求值流程就是尝试调用type(a).__matmul__(a, b)失败后尝试反向方法最终决定是否抛出TypeError。这个行为由 PEP 465 确立至今没有任何变化。1.2 从 numpy 的痛点到 PEP 465在出现之前numpy 用户要算矩阵乘法得写np.dot(a, b)或者后来的np.matmul(a, b)。如果你做过科学计算应该能体会这种表达式的痛苦一个很自然的数学公式C AB在代码里却变成了函数嵌套调用。更麻烦的是np.dot在多维数组上的语义和严格意义的矩阵乘法并不完全一致很多人在这里踩过坑。2015 年Python 社区终于决定结束这场“用函数模拟运算符”的将就PEP 465 正式把引入语言语法。核心动机就一句话矩阵乘法是线性代数里的基本运算理应有一等语法支持。从那时起numpy 里a b就等于np.matmul(a, b)而a * b仍然保持逐元素乘法的语义。两者终于分家了。这里我多说一句为什么不能直接用*numpy 的*在数组语境下代表 Hadamard 积逐元素相乘这是从 MATLAB 时代就延续下来的习惯无数代码依赖这个行为。直接改变*的语义会造成巨大破坏正确做法就是新增一个运算符把矩阵乘法和逐元素乘法彻底分开。2.__matmul__协议完整拆解三个方法一个都不能少2.1 正向方法__matmul__(self, other)正主是__matmul__它定义了self other的运算逻辑。方法签名和__add__等二元运算完全一致接收另一个操作数返回运算结果如果发现自己处理不了这个类型就返回NotImplemented这个特殊值把决定权交给解释器。一个最基本的__matmul__实现逻辑是这样的def __matmul__(self, other): if not isinstance(other, Matrix): return NotImplemented # 具体矩阵乘法逻辑 ... return result注意这里不能直接抛TypeError而是应该返回NotImplemented。因为左侧操作数处理不了不代表右侧操作数也处理不了可能右侧对象有自己的__rmatmul__逆向实现。2.2 反向方法__rmatmul__(self, other)other self这种表达式如果左侧的other没有实现__matmul__或者实现的__matmul__返回了NotImplementedPython 就会尝试调用右侧对象即self的__rmatmul__(other)。这里的参数顺序要注意__rmatmul__(self, other)中self是右边的对象other是左边的对象。这个设计让“混合类型运算”成为可能。举个具体例子如果我用[[1, 2], [3, 4]] my_matrix这种写法左侧是一个普通 listPython 的 list 根本没有__matmul__于是会去右侧的my_matrix.__rmatmul__里寻找机会。我们可以在这个方法里把 list 强转成 Matrix或者直接返回NotImplemented让解释器报错。这就是为什么一个好的矩阵类必须同时实现正反两个方法。2.3 原位方法__imatmul__(self, other)第三个是__imatmul__对应赋值运算。大多数人会忽略这个方法但它的存在与否直接影响对象的身份标识是否改变。如果没有定义__imatmul__执行a b时 Python 会退化成a a b也就是说重新给变量绑定了一个新对象。对不可变对象来说这是唯一正确的方案但对可变对象来说原地更新往往更高效也能保持外部引用的一致性。所以如果你的矩阵类是可变对象且希望m n真正修改m内部的数据就应该自己实现__imatmul__。把三个方法的职责整理一下表达式调用优先级a ba.__matmul__(b)失败则尝试b.__rmatmul__(a)a ba.__imatmul__(b)不存在则退化为a a b3. 手写一个支持的纯 Python 矩阵类3.1 定义数据结构与基础校验光讲协议太虚直接上代码。我用最基础的双层 list 来存储矩阵数据不引入 numpy这样才能真正看清背后的计算逻辑。这个类只做教学演示性能不追求极致但结构要干净。class Matrix: def __init__(self, data): if not isinstance(data, (list, tuple)): raise TypeError(Matrix data must be a 2D list/tuple) rows len(data) if rows 0: raise ValueError(Matrix must have at least one row) cols len(data[0]) if not all(len(row) cols for row in data): raise ValueError(All rows must have the same length) self.data [list(row) for row in data] self.shape (rows, cols) def __repr__(self): return fMatrix({self.data!r}) def __eq__(self, other): if isinstance(other, Matrix): return self.data other.data return NotImplemented校验这一步不是强迫症而是矩阵运算的前提。矩阵乘法要求左矩阵的列数等于右矩阵的行数如果数据本身都不规整后面每算一次都会踩一遍不可预知的坑。shape属性提前存好后面校验直接用不需要每次数一遍行和列。3.2 实现三个魔法方法接下来写核心的__matmul__。标准矩阵乘法的规则是结果矩阵的第 i 行第 j 列元素等于左矩阵第 i 行与右矩阵第 j 列逐项相乘再求和。我用三层循环实现最外层遍历结果的行第二层遍历结果的列最内层做点积累加。def __matmul__(self, other): if not isinstance(other, Matrix): return NotImplemented r1, c1 self.shape r2, c2 other.shape if c1 ! r2: raise ValueError( fMatrix shapes not aligned: {self.shape} {other.shape} ) result [] for i in range(r1): row [] for j in range(c2): total 0 for k in range(c1): total self.data[i][k] * other.data[k][j] row.append(total) result.append(row) return Matrix(result)这个实现虽然朴素但把的计算语义讲得很清楚。如果两个矩阵是 Python 内置的 listlist list会直接报TypeError因为 list 根本没有__matmul__。有了这个 Matrix 类之后才真正有意义。然后是反向方法。这里有个常见的写法误区很多人会把__rmatmul__直接写成return self.__matmul__(other)这其实求的是other self的结果吗不对self.__matmul__(other)算的是self other方向完全反了。正确的逻辑是把左操作数接过来要么转换成 Matrix要么直接复用对方的乘法实现。def __rmatmul__(self, other): if isinstance(other, Matrix): return other.__matmul__(self) if isinstance(other, (list, tuple)): return Matrix(other).__matmul__(self) return NotImplemented最后是__imatmul__。如果左右操作数形状完全一致可以直接原地覆盖否则说明维数不对原地更新没有意义。原地操作可以让m n保持m的对象身份不变这对需要持有矩阵引用的代码很关键。def __imatmul__(self, other): if not isinstance(other, Matrix): other Matrix(other) result self other self.data result.data self.shape result.shape return self3.3 用这个类跑一个二维旋转验证写完类不能空口说能跑直接验证一下。二维旋转矩阵的公式是[cosθ -sinθ] [x] [x] [sinθ cosθ] [y] [y]下面这段代码把点 (1, 0) 旋转 90 度预期结果是 (0, 1)import math theta math.pi / 2 rot Matrix([ [math.cos(theta), -math.sin(theta)], [math.sin(theta), math.cos(theta)], ]) point Matrix([[1], [0]]) rotated rot point print(rotated) # Matrix([[6.123233995736766e-17], [1.0]])结果里的6.12e-17是浮点数误差数学上就是 0。这说明已经正常工作了。你再试试rotated rot [[1], [0]]会出错因为左侧的 rot 是 Matrix右侧是普通 listMatrix 的__matmul__只接受 Matrix。你可以在__matmul__里也加一段 list 转换逻辑通常我会为了兼容性把它补上但为了让协议演示更清晰这里先保留严格的类型判断。4. 实战用__matmul__实现矩阵快速幂4.1 矩阵快速幂的原理快速幂是算法里的经典套路。普通整数快速幂的核心思想是计算a^n时把指数拆成二进制最多只需要log(n)次乘法而不是n次。比如a^13 a^8 * a^4 * a因为13 0b1101。矩阵快速幂完全一样只是把整数乘法替换成矩阵乘法也就是。在实现之前我得先搞定单位矩阵这个概念。单位矩阵就是对角线全为 1、其余全为 0 的方阵任何矩阵乘以单位矩阵都等于它自己。快速幂的初始累乘结果就设为这个单位矩阵才能保证算法正确。如果你自己实现了 Matrix 类还可以顺手加上__pow__方法支持M ** n不过这里为了聚焦我直接写一个独立函数def mat_pow(base, exp): # 返回 base^expbase 必须是支持 运算的对象 n base.shape[0] result Matrix([ [1 if i j else 0 for j in range(n)] for i in range(n) ]) while exp 0: if exp 1: result result base base base base exp 1 return result这个函数只要你传进来的对象支持不管是我们的 Matrix 还是 numpy 数组都能正常工作。这就是鸭子类型的好处——我不关心对象是什么类型只要它有__matmul__一切好说。4.2 不用 numpy 也能算大数斐波那契斐波那契数列有一个经典的矩阵表示[F(n1) F(n) ] [[1, 1], [1, 0]] 的 n 次方 [F(n) F(n-1)]也就是说先构造矩阵M [[1, 1], [1, 0]]计算M^n结果矩阵右上角就是F(n)。传统递归算斐波那契是指数级复杂度动态规划是O(n)而矩阵快速幂能做到O(log n)的时间复杂度。直接跑一下看看def fib(n): M Matrix([[1, 1], [1, 0]]) power mat_pow(M, n) return power.data[0][1] for n in range(10): print(n, fib(n))输出前几项是0, 1, 1, 2, 3, 5, 8, 13, 21, 34完全正确。我再试试fib(1000)它能在瞬间算出一个 209 位的巨大整数因为 Python 的整数是任意精度的不会像 numpy 的固定整数类型那样溢出。这一点是纯 Python 实现的一个隐藏优势。5. 实战用一行写出线性回归闭式解5.1 最小二乘法的矩阵形式在机器学习入门阶段线性回归是最基础的模型。如果用矩阵形式表达假设特征矩阵是X标签向量是y我们要找到权重向量w使得Xw尽量接近y。最小二乘法通过最小化损失函数||Xw - y||^2求导后解出来的闭式解是w (X^T X)^(-1) X^T y这个公式写成长 Python 代码需要 3 次矩阵乘法、1 次转置、1 次求逆如果在没有的旧时代你得写成np.linalg.inv(np.dot(np.dot(X.T, X), X.T)).dot(y)这种嵌套地狱。有了之后公式和代码几乎一一对应。5.2 numpy 版本实现与对比我直接用 numpy 演示因为它是科学计算的标配环境而且 numpy 数组的正好走的就是__matmul__协议底层是优化过的 BLAS 实现。import numpy as np X np.array([ [1, 1], [1, 2], [1, 3], [1, 4], ]) y np.array([2, 4, 6, 8]) theta np.linalg.inv(X.T X) X.T y print(theta) # [0. 2.]结果[0, 2]完美拟合了y 2x这条直线。注意这里第一列全为 1代表截距项第二种列为真实特征。计算中用了两处X.T X得到参数矩阵(... ) X.T y得到最终权重。这里有个细节值得新人注意X.T X和X.T * X是完全不同的。前者是矩阵乘法结果是一个 2x2 矩阵后者会被 numpy 当作逐元素乘法尝试广播后直接报错或得到错误结果。的存在让这两者的区分变得极其自然不再依赖np.dot这种函数名去猜了。为了更直观地看到的价值我把类比的例子放在表格里传统写法现代写法语义np.dot(a, b)a b矩阵乘法np.multiply(a, b)a * b逐元素乘法np.linalg.inv(a).dot(b)np.linalg.inv(a) b解线性系统你会发现用之后代码的可读性提升了不止一个档次。数学公式长什么样代码就长什么样不需要在脑内反复做“公式到函数名”的映射。6. 避坑指南与 Python 3.12 带来的变化6.1 六个高频问题速查我在设计和调试矩阵类的过程中踩过不少坑这里整理成速查表每一个都是真实会遇到的。问题原因解决方案TypeError: unsupported operand type(s) for 操作数没有实现__matmul__比如两个 list 直接先转成 Matrix 或 numpy 数组虚拟环境场景下之后对象身份变了没有实现__imatmul__解释器退化为a a b实现__imatmul__原地更新 dataother self结果方向错了__rmatmul__里错误调用self.__matmul__(other)应该调用other.__matmul__(self)或在转换后让other主导运算魔法方法返回了错误类型__matmul__直接返回了NotImplemented以外的东西比如None保证运算结果始终是 Matrix 或标量矩阵维度不匹配但报错信息不清晰没有在乘法前做形状校验在__matmul__开头检查self.shape[1] other.shape[0]忘记定义__rmatmul__导致混合类型运算全挂只实现了正向方法反向没有兜底始终同时实现__rmatmul__和__matmul__这里重点标一下维度校验的体验问题。如果你不在__matmul__里提前检查形状Python 的三层循环会在某个k下标越界时报出莫名其妙的IndexError。一个清晰的ValueError消息能帮你省掉十分钟的排查时间这就是为什么矩阵类里一定要有 shape 记录。6.2 Python 3.12 对魔方方法查找做了底层优化很多人关心 Python 3.12 对__matmul__有什么影响。结论是语法和协议完全没变但底层实现有优化。CPython 在 3.12 里改进了特殊方法包括__matmul__的隐式查找机制。以前每次执行a b时解释器都要通过类型对象的 MRO 链去查找__matmul__是否存在3.12 之后这个查找结果会被缓存在类型对象里后续调用走缓存路径减少了大量属性查找的开销。对于科学计算场景里高频执行的运算这种底层优化是实实在在的收益而且你不需要改任何代码。还有一个和 3.12 相关的点Python 3.12 开始提供了更清晰的方法属性检查方式比如inspect.getattr_static之类的工具让开发者能更安全地探索类型属性。但说实话日常开发和运算关系不大。只要记住一点你的代码如果能在 3.10、3.11 上正常使用放到 3.12 上不会有任何问题甚至可能更快。6.3 性能分析与设计心得纯 Python 写的 Matrix 类矩阵乘法的时间复杂度是O(n^3)。如果两个 100x100 的矩阵相乘最内层循环要跑一百万次每次都是 Python 解释器级别的循环和列表索引速度会非常感人。我实际测下来大约需要几百毫秒到一秒而 numpy 用底层优化过的 BLAS同样的运算只需要几毫秒甚至更低性能差距可能在百倍以上。但这不是说纯 Python 的实现没有意义。教学价值自不必说当你亲手用三层循环写出矩阵乘法你就彻底理解到底在算什么不会再把*和搞混。另外在某些不允许引入 numpy 的环境中自己写一个只支持小规模运算的 Matrix 类也能解决实际需求。设计上我还有几个心得体会尽量保持的语义一致别人看到你的类支持默认期望是“矩阵乘法”或“线性变换组合”不要做太另类的实现。善用NotImplemented而不是直接抛异常这能让混合类型运算有回旋余地给右侧操作数一个响应的机会。如果你在自定义类里用到务必想清楚对象的“身份”id要不要保持不变。不可变对象不需要__imatmul__可变对象建议实现。我个人在实际操作中的体会是真正理解一个运算符重载最好的方式不是读文档而是亲手把它的正向、反向、原位三个方法都实现一遍再写两个实战项目验证。等你发现自己的类能被无缝接入算法代码、和 numpy 数组混用都不违和的那一刻__matmul__才是真正的“会了”。最后再分享一个小技巧你完全可以把__matmul__用作一种“领域特定语言”的载体——比如定义一个变换链类让A B C读起来就像数学论文里的变换组合或者把当成语义化的复合操作符只要你能说清楚这个运算的代数意义。灵活使用但务必保持语义可预期这是所有魔术方法设计的底线。