如何参数化元组类型提示中np.ndarray的返回值数量?
问题描述
假设有如下带类型提示的代码:
import numpy as np def func() -> tuple[np.ndarray, np.ndarray, np.ndarray]: arr1 = np.empty(shape=(5,)) arr2 = np.ones(shape=(5,)) arr3 = np.zeros(shape=(5,)) return arr1, arr2, arr3
当返回的np.ndarray数量增多时,手动逐个编写类型提示会非常繁琐。有没有办法无需逐个编写,就能为指定数量的np.ndarray添加类型提示?比如类似这种(无法运行的示例写法):
import numpy as np def func() -> tuple[*([np.ndarray]*3)]: arr1 = np.empty(shape=(5,)) arr2 = np.ones(shape=(5,)) arr3 = np.zeros(shape=(5,)) return arr1, arr2, arr3
解决方案
1. Python 3.11+:用Unpack实现动态展开(最接近需求)
Python 3.11引入了Unpack类型,可以配合元组乘法生成指定数量的重复类型,写法完全贴合你的需求:
import numpy as np from typing import Unpack def func() -> tuple[Unpack[tuple[np.ndarray]*3]]: arr1 = np.empty(shape=(5,)) arr2 = np.ones(shape=(5,)) arr3 = np.zeros(shape=(5,)) return arr1, arr2, arr3
这里tuple[np.ndarray]*3生成包含3个np.ndarray的类型元组,再通过Unpack展开到返回值的tuple类型中,主流类型检查器(如mypy、pyright)均支持该写法。
2. 类型别名简化(兼容旧版本Python)
如果需要兼容Python 3.11之前的版本,可以定义类型别名复用重复的类型提示,避免重复编写:
import numpy as np from typing import Tuple # 定义包含3个np.ndarray的元组类型别名 TripleNDArray = Tuple[np.ndarray, np.ndarray, np.ndarray] def func() -> TripleNDArray: arr1 = np.empty(shape=(5,)) arr2 = np.ones(shape=(5,)) arr3 = np.zeros(shape=(5,)) return arr1, arr2, arr3
后续需要更多数量时,只需修改别名定义即可,比如5个元素就写FiveNDArray = Tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray, np.ndarray]。
3. 任意数量的简化写法(无需固定数量)
如果不需要严格限制返回值的元素数量,只需所有元素都是np.ndarray,可以用可变长度元组的类型提示:
import numpy as np def func() -> tuple[np.ndarray, ...]: arr1 = np.empty(shape=(5,)) arr2 = np.ones(shape=(5,)) arr3 = np.zeros(shape=(5,)) return arr1, arr2, arr3
这种写法表示返回一个由任意数量np.ndarray组成的元组,类型检查器会接受所有符合元素类型的元组。
内容的提问来源于stack exchange,提问作者Rocajoy
相关产品推荐
相关产品推荐

