Python 3.12 Magic Method – __matmul__(self, other)
__matmul__ 是 Python 中用于定义 矩阵乘法运算符 @ 的核心魔术方法。该运算符是 Python 3.5 通过 PEP 465 引入的,旨在为数值计算提供一个专用的中缀矩阵乘法运算符,从而解决长期以来 * 运算符在元素乘法和矩阵乘法之间的语义冲突 。正确实现 __matmul__ 可以让自定义类(如矩阵、向量、线性变换)支持 @ 运算,并与 NumPy 等科学计算库无缝协作。本文将详细解析其定义、底层机制、设计原则,并通过多个示例逐行演示如何正确实现。
1. 定义与签名
def __matmul__(self, other) –> object:
...
- 参数:
- self:当前对象(左操作数)。
- other:另一个操作数(右操作数),可以是任意类型。
- 返回值:应返回一个新的对象,代表矩阵乘法的结果。如果运算未定义(例如类型不兼容或维度不匹配),应返回单例 NotImplemented。
- 调用时机:
- x @ y 会首先尝试调用 x.__matmul__(y) 。
- 如果 x.__matmul__(y) 返回 NotImplemented,Python 会尝试调用 y.__rmatmul__(x)(反向矩阵乘法)。
- 如果两者都返回 NotImplemented,最终抛出 TypeError 。
2. 用途与典型场景
- 矩阵乘法:自定义矩阵类实现 @ 进行数学上的矩阵乘积 。
- 向量点积:虽然 @ 的名称是“矩阵乘法”,但也可用于实现向量的点积(1D 数组)。
- 线性变换:将变换矩阵应用于向量 。
- 与 NumPy 集成:NumPy 的 ndarray 实现了 __matmul__,使得 A @ B 成为推荐的矩阵乘法语法 。
- 自定义运算:虽然 @ 的设计初衷是矩阵乘法,但它本质上只是一个运算符,可以赋予任何含义,例如计算两点间距离 。
矩阵乘法通常不满足交换律(除非在特殊情况下),因此 __rmatmul__ 的实现通常需要独立处理,不能简单地委托给 __matmul__。
3. 底层实现机制
在 Python/C API 层面,每个类型对象(PyTypeObject)都有一个 tp_as_number 结构体,其中包含 nb_matrix_multiply 槽位,这是一个函数指针,用于处理矩阵乘法操作 。当执行 x @ y 时,解释器会遵循标准的反向运算符查找机制:
对于 Python 层定义的 __matmul__,它会被包装到 nb_matrix_multiply 槽位中。反向方法 __rmatmul__ 也会在必要时被调用,用于处理左操作数不支持该运算的情况 。
4. 设计原则与最佳实践
- 遵循数学规则:对于矩阵乘法,应严格遵循维度兼容性规则:第一个矩阵的列数必须等于第二个矩阵的行数。维度不匹配时应抛出 ValueError 。
- 返回新对象:__matmul__ 通常不应修改操作数本身,而应返回一个包含运算结果的新对象。这符合数学运算的不可变习惯。
- 类型检查:应检查 other 的类型是否兼容,如果类型不匹配,应返回 NotImplemented,而不是抛出异常 。这样给另一操作数提供尝试反向运算的机会。
- 实现反向方法:为了支持混合类型运算(如 int @ MyMatrix),应实现 __rmatmul__ 。
- 与 __imatmul__ 的区分:__imatmul__ 用于就地矩阵乘法(@=),通常应修改自身并返回 self,适用于可变对象 。
- 性能考虑:对于大型矩阵,应使用高效的算法(如 Strassen 算法)或直接委托给 NumPy 等优化库。
5. 示例与逐行解析
示例 1:基本矩阵类(实现矩阵乘法)
class Matrix:
def __init__(self, data):
"""初始化矩阵,data 为二维列表"""
self.data = data
self.rows = len(data)
self.cols = len(data[0]) if data else 0
def __matmul__(self, other):
"""实现矩阵乘法 A @ B"""
# 1. 类型检查
if not isinstance(other, Matrix):
return NotImplemented
# 2. 维度兼容性检查
if self.cols != other.rows:
raise ValueError(f"Incompatible dimensions: {self.rows}x{self.cols} and {other.rows}x{other.cols}")
# 3. 计算结果矩阵(大小为 self.rows x other.cols)
result = [[0 for _ in range(other.cols)] for _ in range(self.rows)]
# 4. 三重循环计算矩阵乘法
for i in range(self.rows):
for j in range(other.cols):
for k in range(self.cols):
result[i][j] += self.data[i][k] * other.data[k][j]
# 5. 返回新的 Matrix 对象
return Matrix(result)
def __repr__(self):
return '\\n'.join([' '.join(map(str, row)) for row in self.data])
逐行解析 :
| 1-5 | __init__ | 存储矩阵数据,记录行数和列数。 |
| 6-20 | __matmul__ | 定义矩阵乘法。 |
| 7-9 | 类型检查 | 如果 other 不是 Matrix,返回 NotImplemented,让 Python 尝试反向运算。 |
| 11-13 | 维度检查 | 确保 self.cols == other.rows,否则抛出 ValueError,符合矩阵乘法规则。 |
| 15-16 | 初始化结果矩阵 | 创建大小为 self.rows x other.cols 的零矩阵。 |
| 17-20 | 三重循环 | 标准矩阵乘法算法:result[i][j] = sum(self.data[i][k] * other.data[k][j] for k in range(self.cols))。 |
| 22 | 返回新对象 | 用计算结果创建新的 Matrix 实例。 |
为什么这样写?
- 严格遵循矩阵乘法的数学定义和维度检查,确保运算的正确性。
- 返回新对象,保持不可变性,避免意外修改原矩阵。
- 类型检查后返回 NotImplemented,使 Python 有机会尝试反向调用(如 int @ Matrix)。
验证:
A = Matrix([[1, 2], [3, 4]])
B = Matrix([[5, 6], [7, 8]])
print("A @ B 结果:")
print(A @ B) # 调用 A.__matmul__(B)
运行结果:
A @ B 结果:
19 22
43 50
示例 2:实现反向矩阵乘法(支持混合类型)
当左操作数为内置类型(如 int)时,需要实现 __rmatmul__ 来处理。
class Matrix:
def __init__(self, data):
self.data = data
self.rows = len(data)
self.cols = len(data[0]) if data else 0
def __matmul__(self, other):
"""实现矩阵乘法 A @ B"""
# 1. 类型检查
if not isinstance(other, Matrix):
return NotImplemented
# 2. 维度兼容性检查
if self.cols != other.rows:
raise ValueError(f"Incompatible dimensions: {self.rows}x{self.cols} and {other.rows}x{other.cols}")
# 3. 计算结果矩阵(大小为 self.rows x other.cols)
result = [[0 for _ in range(other.cols)] for _ in range(self.rows)]
# 4. 三重循环计算矩阵乘法
for i in range(self.rows):
for j in range(other.cols):
for k in range(self.cols):
result[i][j] += self.data[i][k] * other.data[k][j]
# 5. 返回新的 Matrix 对象
return Matrix(result)
def __rmatmul__(self, other):
# 处理 other @ Matrix 的情况
if isinstance(other, (int, float)):
# 标量与矩阵乘法:将标量视为对角矩阵?或标量乘?通常标量乘定义为数乘
# 这里我们定义为数乘(每个元素乘以标量)
result = [[other * val for val in row] for row in self.data]
return Matrix(result)
return NotImplemented
def __repr__(self):
return '\\n'.join([' '.join(map(str, row)) for row in self.data])
解析 :
- 2 @ M 首先尝试 int.__matmul__(M),但 int 没有实现 __matmul__,返回 NotImplemented。
- 然后尝试 M.__rmatmul__(2),我们定义了数乘运算,返回新的矩阵。
- 注意:这里我们定义了标量左乘矩阵为数乘(每个元素乘以标量),这是一种合理的解释,但不是唯一的(也可以定义为其他运算)。开发者应根据业务逻辑明确定义。
验证:
M = Matrix([[1, 2], [3, 4]])
result = 2 @ M # 调用 M.__rmatmul__(2)
print(result)
运行结果:
2 4
6 8
示例 3:向量点积
@ 也可用于实现向量的点积,这是矩阵乘法在一维数组上的特例 。
class Vector:
def __init__(self, components):
self.components = components
def __matmul__(self, other):
if not isinstance(other, Vector):
return NotImplemented
if len(self.components) != len(other.components):
raise ValueError("Vectors must have same length")
# 计算点积:sum(x_i * y_i)
return sum(a * b for a, b in zip(self.components, other.components))
def __repr__(self):
return f"Vector({self.components})"
解析: 点积结果是一个标量,而不是 Vector 对象,这与矩阵乘法返回矩阵不同,但完全符合数学定义。
验证:
v1 = Vector([1, 2, 3])
v2 = Vector([4, 5, 6])
print(v1 @ v2) # 32 (1*4 + 2*5 + 3*6)
运行结果:
32
示例 4:就地矩阵乘法(@=)
对于可变对象,可以实现 __imatmul__ 支持就地修改 。
class Matrix:
def __init__(self, data):
self.data = data
self.rows = len(data)
self.cols = len(data[0]) if data else 0
def __imatmul__(self, other):
# 就地矩阵乘法:self = self @ other
if not isinstance(other, Matrix):
return NotImplemented
if self.cols != other.rows:
raise ValueError("Incompatible dimensions")
# 计算结果(创建临时矩阵,避免在计算过程中修改 self.data)
result = [[0 for _ in range(other.cols)] for _ in range(self.rows)]
for i in range(self.rows):
for j in range(other.cols):
for k in range(self.cols):
result[i][j] += self.data[i][k] * other.data[k][j]
# 更新自身
self.data = result
self.rows = len(result)
self.cols = len(result[0]) if result else 0
return self
def __repr__(self):
return '\\n'.join([' '.join(map(str, row)) for row in self.data])
解析: __imatmul__ 必须返回 self,且通常应修改对象自身。这里我们先计算结果的临时矩阵,然后更新 self.data,避免在计算过程中修改原数据。
验证:
M = Matrix([[1, 2], [3, 4]])
M @= Matrix([[2, 0], [1, 2]]) # 调用 M.__imatmul__()
print(M)
运行结果:
4 4
10 8
示例 5:自定义运算——计算两点间距离
虽然 @ 的设计初衷是矩阵乘法,但 Python 允许你赋予它任何含义。以下示例使用 @ 计算平面上两点间的欧几里得距离 。
from math import sqrt
class Point:
def __init__(self, x, y):
self.x = x
self.y = y
def __matmul__(self, other):
if not isinstance(other, Point):
return NotImplemented
# 计算欧几里得距离
return sqrt((self.x – other.x) ** 2 + (self.y – other.y) ** 2)
def __repr__(self):
return f"Point({self.x}, {self.y})"
解析: 这个例子展示了运算符重载的灵活性,但需谨慎使用,避免与常规语义混淆。对于代码的可读性,通常建议保持 @ 的矩阵乘法语义,除非有非常明确的理由。
验证:
a = Point(1, 3)
b = Point(4, 7)
print(a @ b) # 5.0
运行结果:
5.0
6. 与 __rmatmul__ 和 __imatmul__ 的关系
| __matmul__(self, other) | 正向矩阵乘法 self @ other | 新对象 | x @ y |
| __rmatmul__(self, other) | 反向矩阵乘法 other @ self | 新对象 | 正向返回 NotImplemented 时 |
| __imatmul__(self, other) | 就地矩阵乘法 self @= other | self | x @= y |
关键区别:
- __rmatmul__ 用于左操作数不支持运算时的反向调用,由于矩阵乘法不满足交换律,必须独立实现。
- __imatmul__ 用于原地修改对象,应返回 self,适用于可变对象。
7. 注意事项与陷阱
- 不要修改 self:__matmul__ 应返回新对象,除非类是可变的且你明确希望就地修改(此时应使用 __imatmul__)。
- 正确使用 NotImplemented:当类型不兼容时返回 NotImplemented,而不是抛出异常 。这给另一侧机会处理。
- 维度检查:矩阵乘法必须检查维度兼容性,不匹配时应抛出 ValueError 。
- 与 NumPy 的交互:如果使用 NumPy,numpy.ndarray 已经实现了 __matmul__,可以直接使用 A @ B 。
- 避免滥用:虽然可以自定义 @ 的任何含义,但为了代码可读性,建议遵循其矩阵乘法的设计初衷。
8. 总结
| 角色 | 定义矩阵乘法运算符 @ |
| 签名 | __matmul__(self, other) -> object |
| 返回值 | 新对象,或 NotImplemented |
| 调用时机 | x @ y,以及反向尝试 |
| 底层 | C 层的 nb_matrix_multiply 槽位 |
| 与 __rmatmul__ 的关系 | 反向矩阵乘法,用于 other @ self,需独立实现 |
| 与 __imatmul__ 的关系 | 就地矩阵乘法 @=,通常修改自身并返回 self |
| 最佳实践 | 返回新对象、维度检查、类型检查、使用 NotImplemented |
掌握 __matmul__ 是实现自定义数值类型(特别是线性代数相关类)的重要一环。通过理解其底层机制和设计原则,你可以构建出与 NumPy 等科学计算库无缝协作的类,同时保持代码的数学清晰性。
如果在学习过程中遇到问题,欢迎在评论区留言讨论!

