Python中支持@矩阵乘法运算的对象该如何正确标注类型?
支持@矩阵乘法运算的实体的通用类型标注方案
Python 中@运算符的底层对应__matmul__魔法方法,要兼容numpy数组、scipy稀疏矩阵等所有支持该运算的类型,最灵活的方案是使用typing.Protocol定义结构化类型协议,不需要绑定具体实现类:
- 先定义泛型协议约束矩阵乘法的输入输出类型:
from typing import Protocol, TypeVar import numpy as np from numpy.typing import ArrayLike # Python 3.8以下版本需先安装typing_extensions,再从typing_extensions导入Protocol MatMulResult = TypeVar("MatMulResult") class SupportsMatMul(Protocol[MatMulResult]): def __matmul__(self, other: ArrayLike) -> MatMulResult: ...
- 给
fun函数添加类型标注:
def fun(A: SupportsMatMul[np.ndarray], x: ArrayLike) -> np.ndarray: return (A @ x) ** 2 - 27.0
以上标注采用结构化匹配逻辑,只要传入的A实现了符合签名的__matmul__方法,且和x运算后返回numpy数组,就能通过类型检查,不需要额外适配numpy、scipy等库的具体矩阵类型。
如果你的使用场景不需要支持自定义矩阵类型,也可以直接用联合类型标注:Union[np.ndarray, scipy.sparse.spmatrix],不过扩展性较差,新增支持矩阵乘法的类型时需要手动更新联合类型列表。
内容的提问来源于stack exchange,提问作者Nico Schlömer
相关产品推荐
相关产品推荐

