为scipy.stats中的通用/冻结多元分布添加合规的类型提示
scipy.stats中的通用/冻结多元分布添加合规的类型提示
我明白你遇到的问题了——scipy的多元分布类型注解确实不如单变量的完善,mypy没法正确识别multi_rv_frozen的dim、rvs属性,也不认为multi_rv_generic是可调用的。而且依赖scipy.stats._multivariate里的私有类型也不是长期可靠的做法。下面给你几个靠谱的解决思路:
方案1:用Protocol定义接口(推荐,类型安全且不依赖私有模块)
Protocol是Python类型系统里的“鸭子类型”声明,我们可以自己定义多元分布需要满足的接口,不管scipy内部怎么实现,只要对象符合这个接口,mypy就能正确识别。这个方法既能摆脱对私有模块的依赖,又能保持严格的类型检查。
完整可运行且mypy无报错的代码示例:
from typing import Any, Protocol, Callable, TypeVar # 定义冻结多元分布的接口:必须具备dim属性和rvs采样方法 class FrozenMultivariateDist(Protocol): @property def dim(self) -> int: ... # 仅做类型声明,无需实现 def rvs(self, n_samples: int, **kwargs: Any) -> Any: ... # 仅做类型声明,无需实现 # 定义通用多元分布的接口:必须是可调用的,调用后返回冻结的多元分布 FrozenDistT = TypeVar('FrozenDistT', bound=FrozenMultivariateDist) class GenericMultivariateDist(Protocol, Callable[..., FrozenDistT]): ... # 仅做类型声明,无需实现 def sample_frozen_multivariate(n_samples: int, n_variates: int, dist: FrozenMultivariateDist): if dist.dim != n_variates: msg = 'distribution dimension %s inconsistent with n_variates=%s' raise ValueError(msg % (dist.dim, n_variates)) sample = dist.rvs(n_samples) return sample def sample_generic_multivariate(n_samples: int, n_variates: int, dist: GenericMultivariateDist, *distparams: Any): frozen_dist = dist(*distparams) if frozen_dist.dim != n_variates: msg = 'distribution dimension %s inconsistent with n_variates=%s' raise ValueError(msg % (frozen_dist.dim, n_variates)) sample = frozen_dist.rvs(n_samples) return sample # 测试代码 if __name__ == "__main__": from scipy.stats import multivariate_normal # 测试冻结的多元分布 n_samples = 4 n_variates = 2 frozen_dist = multivariate_normal(mean=[0, 0], cov=[[1, 0], [0, 1]]) print(sample_frozen_multivariate(n_samples, n_variates, frozen_dist)) # 测试通用的多元分布 n_samples = 4 n_variates = 2 generic_dist = multivariate_normal mean, cov = [-1, 1], [[1, 0], [0, 1]] print(sample_generic_multivariate(n_samples, n_variates, generic_dist, mean, cov))
方案2:临时抑制mypy错误(快速解决,但不推荐)
如果你只是想快速让mypy通过,不想大幅修改代码,可以用cast或者# type: ignore来抑制错误,但这种方法会丢失部分类型检查能力,而且仍然依赖scipy的私有模块,不适合长期项目:
from typing import Any, Callable, cast from scipy.stats import multivariate_normal from scipy.stats._multivariate import multi_rv_generic, multi_rv_frozen def sample_frozen_multivariate(n_samples: int, n_variates: int, dist: multi_rv_frozen): # 用cast告诉mypy dist具备dim属性和rvs方法 dist_typed = cast(Any, dist) if dist_typed.dim != n_variates: msg = 'distribution dimension %s inconsistent with n_variates=%s' raise ValueError(msg % (dist_typed.dim, n_variates)) sample = dist_typed.rvs(n_samples) return sample def sample_generic_multivariate(n_samples: int, n_variates: int, dist: multi_rv_generic, *distparams: Any): # 用cast告诉mypy dist是可调用的,返回multi_rv_frozen dist_callable = cast(Callable[..., multi_rv_frozen], dist) frozen_dist = dist_callable(*distparams) frozen_dist_typed = cast(Any, frozen_dist) if frozen_dist_typed.dim != n_variates: msg = 'distribution dimension %s inconsistent with n_variates=%s' raise ValueError(msg % (frozen_dist_typed.dim, n_variates)) sample = frozen_dist_typed.rvs(n_samples) return sample # 测试代码 if __name__ == "__main__": n_samples = 4 n_variates = 2 frozen_dist = multivariate_normal(mean=[0, 0], cov=[[1, 0], [0, 1]]) print(sample_frozen_multivariate(n_samples, n_variates, frozen_dist)) n_samples = 4 n_variates = 2 generic_dist = multivariate_normal mean, cov = [-1, 1], [[1, 0], [0, 1]] print(sample_generic_multivariate(n_samples, n_variates, generic_dist, mean, cov))
方案3:补充scipy的类型存根(进阶)
如果你想长期依赖scipy的内部类型,可以给mypy添加自定义的类型存根文件(.pyi),补充multi_rv_generic和multi_rv_frozen的类型注解。不过这个方法需要维护存根文件,适配scipy的版本更新,适合对类型检查要求极高的进阶场景。
备注:内容来源于stack exchange,提问作者brentertainer
相关产品推荐
相关产品推荐

