欢迎光临
我们一直在努力

Python 3.12 MagicMethods - 47 - __matmul__

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 时,解释器会遵循标准的反向运算符查找机制:

  • 获取 x 的类型对象的 tp_as_number 结构。
  • 如果存在 nb_matrix_multiply,则调用它,传入 x 和 y,返回结果对象或 Py_NotImplemented。
  • 如果 x 的 nb_matrix_multiply 返回 Py_NotImplemented,则尝试获取 y 的类型对象的 nb_matrix_multiply,并调用它,但此时参数顺序已交换(即调用 y 的 __rmatmul__ 对应的 C 函数)。
  • 如果仍然失败,则抛出 TypeError。
  • 对于 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 等科学计算库无缝协作的类,同时保持代码的数学清晰性。

    如果在学习过程中遇到问题,欢迎在评论区留言讨论!

    赞(0)
    未经允许不得转载:171主机测评 » Python 3.12 MagicMethods - 47 - __matmul__
    分享到: 更多 (0)

    评论 抢沙发

    • 昵称 (必填)
    • 邮箱 (必填)
    • 网址