如何让支持多类型数值(含NumPy数组)的泛型mean函数通过Mypy类型检查
我来帮你搞定这个泛型mean函数的Mypy类型检查问题。核心痛点在于Mypy没法自动推断泛型类型T支持哪些操作,也没法确认输入的可迭代对象是否能获取长度。我们可以通过定义Protocol来约束类型行为,同时调整函数实现来满足类型检查要求。
最终可通过Mypy检查的代码
from collections import deque from typing import Protocol, TypeVar, Iterable, Sized import numpy as np # 定义一个Protocol,约束支持数值操作的类型 class Numeric(Protocol): def __add__(self, other: 'Numeric') -> 'Numeric': ... def __mul__(self, scalar: float) -> 'Numeric': ... # 绑定泛型T到Numeric Protocol T = TypeVar('T', bound=Numeric) def mean(a: Iterable[T] & Sized) -> T: items = list(a) if not items: raise ValueError("无法计算空可迭代对象的平均值") # 手动累加,避免sum默认初始值0带来的类型不兼容问题 total = items[0] for item in items[1:]: total = total + item # 乘以倒数计算平均值 return total * (1.0 / len(items)) # 测试用例 c = mean([1.0, 1.5]) print(c) a = np.array([1, 2, 3]) b = np.array([4, 5, 6]) c = mean([a, b]) print(c) print(mean((1,2,3))) d = deque([a,b]) print(mean(d))
关键修改点说明
定义
NumericProtocol
这个Protocol相当于给Mypy立下规则:任何能作为T的类型,必须支持加法(__add__)和与浮点数相乘(__mul__)的操作。这样Mypy就知道泛型T可以安全执行这些运算,解决了Unsupported operand types for +/*的错误。约束输入类型为
Iterable[T] & Sized
原代码里len(a)要求输入对象是Sized类型(有长度属性),但普通的Iterable不一定满足(比如生成器就没有长度)。通过Iterable[T] & Sized,我们明确要求输入同时具备可迭代和可获取长度的特性,解决了Argument 1 to "len" has incompatible type的错误。手动代替
sum进行累加
原生sum函数默认初始值是0,Mypy没法推断0和泛型T(比如NumPy数组)的兼容性。手动从第一个元素开始累加,既避开了类型不匹配的坑,也让运算逻辑更直观,解决了Argument 1 to "sum" has incompatible type的错误。自动适配多类型返回值
现在Mypy会根据输入的类型自动推断返回值类型:比如输入float列表返回float,输入NumPy数组列表返回ndarray,解决了Incompatible types in assignment的错误。
额外说明
如果想保留sum函数的使用,可以给sum指定与T匹配的初始值,但需要扩展Protocol来支持获取类型的零值,不过手动累加的方式更简单直接,也更适配多类型场景。
备注:内容来源于stack exchange,提问作者mirekh_68

