Pyright无法推断多态Python函数数据类型问题求助
问题描述
我编写了一个多态Python函数,接收值类型为float、int或numpy.ndarray的字典,内部通过两个子函数分别处理标量(float/int)和数组(np.ndarray)类型,使用ArrayLike进行类型标注。但Pyright无法识别all(isinstance(val, Union[float,int]))这类类型守卫,按照提示改用Mapping后仍出现类型不兼容错误,希望得到解决建议。
初始代码
from typing import Dict, Union, Mapping import numpy as np from numpy.typing import ArrayLike def fun_for_float(data: dict[str, Union[float,int]]) -> None: pass def fun_for_array(data: dict[str, np.ndarray]) -> None: pass def main_function(data: dict[str, ArrayLike]) -> None: if all(isinstance(val, Union[float,int]) for val in data.values()): return fun_for_float(data) elif all(isinstance(val, np.ndarray) for val in data.values()): return fun_for_array(data) else: raise ValueError("Values in the dictionary must be float or np.ndarray.")
初始报错
Argument of type "dict[str, ArrayLike]" cannot be assigned to parameter "data" of type "dict[str, float | int]" in function "fun_for_float" "dict[str, ArrayLike]" is incompatible with "dict[str, float | int]" Type parameter "_VT@dict" is invariant, but "ArrayLike" is not the same as "float | int" Consider switching from "dict" to "Mapping" which is covariant in the value type
修改后代码
def fun_for_float(data: Mapping[str, Union[float,int]]) -> None: pass def fun_for_array(data: Mapping[str, np.ndarray]) -> None: pass def main_function(data: Mapping[str, ArrayLike]) -> None: if all(isinstance(val, Union[float,int]) for val in data.values()): return fun_for_float(data) elif all(isinstance(val, np.ndarray) for val in data.values()): return fun_for_array(data) else: raise ValueError("Values in the dictionary must be float or np.ndarray.")
修改后报错
Argument of type "Mapping[str, ArrayLike]" cannot be assigned to parameter "data" of type "Mapping[str, float | int]" in function "fun_for_float" "Mapping[str, ArrayLike]" is incompatible with "Mapping[str, float | int]" Type parameter "_VT_co@Mapping" is covariant, but "ArrayLike" is not a subtype of "float | int"
解决建议
方法1:使用类型断言(cast)
直接用typing.cast明确告诉类型检查器当前data的具体类型,这是最直接的解决方案:
from typing import Dict, Union, Mapping, cast import numpy as np from numpy.typing import ArrayLike def fun_for_float(data: Mapping[str, Union[float,int]]) -> None: pass def fun_for_array(data: Mapping[str, np.ndarray]) -> None: pass def main_function(data: Mapping[str, ArrayLike]) -> None: # 注意:isinstance不接受Union类型作为第二个参数,需改用元组形式 if all(isinstance(val, (float, int)) for val in data.values()): return fun_for_float(cast(Mapping[str, Union[float, int]], data)) elif all(isinstance(val, np.ndarray) for val in data.values()): return fun_for_array(cast(Mapping[str, np.ndarray], data)) else: raise ValueError("Values in the dictionary must be float or np.ndarray.")
方法2:自定义类型守卫函数
为映射类型编写自定义类型守卫,让Pyright能正确识别你的类型判断逻辑:
from typing import Dict, Union, Mapping, TypeGuard import numpy as np from numpy.typing import ArrayLike def is_float_int_mapping(data: Mapping[str, ArrayLike]) -> TypeGuard[Mapping[str, Union[float, int]]]: return all(isinstance(val, (float, int)) for val in data.values()) def is_ndarray_mapping(data: Mapping[str, ArrayLike]) -> TypeGuard[Mapping[str, np.ndarray]]: return all(isinstance(val, np.ndarray) for val in data.values()) def fun_for_float(data: Mapping[str, Union[float,int]]) -> None: pass def fun_for_array(data: Mapping[str, np.ndarray]) -> None: pass def main_function(data: Mapping[str, ArrayLike]) -> None: if is_float_int_mapping(data): return fun_for_float(data) elif is_ndarray_mapping(data): return fun_for_array(data) else: raise ValueError("Values in the dictionary must be float or np.ndarray.")
自定义类型守卫通过TypeGuard标注返回值,明确告知类型检查器:当函数返回True时,输入参数的类型就是TypeGuard中指定的类型。
方法3:缩小输入类型范围
ArrayLike是一个宽泛的类型(包含列表等可转数组的类型),如果你的函数仅接受float/int或np.ndarray,可以直接限定输入类型,避免宽泛类型带来的兼容性问题:
from typing import Dict, Union, Mapping, cast import numpy as np def fun_for_float(data: Mapping[str, Union[float,int]]) -> None: pass def fun_for_array(data: Mapping[str, np.ndarray]) -> None: pass # 直接将参数类型改为明确的联合类型 def main_function(data: Mapping[str, Union[float, int, np.ndarray]]) -> None: if all(isinstance(val, (float, int)) for val in data.values()): return fun_for_float(cast(Mapping[str, Union[float, int]], data)) elif all(isinstance(val, np.ndarray) for val in data.values()): return fun_for_array(cast(Mapping[str, np.ndarray], data)) else: raise ValueError("Values in the dictionary must be float or np.ndarray.")
额外注意点
- 运行时
Union类型不存在,isinstance(val, Union[float,int])会报错,必须替换为isinstance(val, (float, int))的元组形式。
内容的提问来源于stack exchange,提问作者Holan
相关产品推荐
相关产品推荐

