如何简化Python函数中复杂的返回类型定义
简化复杂Python函数返回类型的两种方案
你可以通过类型别名复用或轻量数据类两种方式,不用Pydantic就能让复杂的返回类型更清晰,同时支持多函数调用复用。
方案1:用类型别名拆分复杂结构(兼容原返回格式)
将重复的类型定义抽成单独的类型别名,放在独立文件中供多函数调用,既保留原返回结构,又让类型提示一目了然。
步骤1:创建共享类型文件(比如types.py)
import numpy as np from typing import TypeAlias, Union, Tuple, List # 单组训练/测试索引的类型 TrainTestIndexes: TypeAlias = Tuple[np.ndarray, np.ndarray] # 带名称的折索引(名称 + 训练测试索引) NamedFold: TypeAlias = Tuple[str | int, TrainTestIndexes] # 最终返回类型:二选一的列表结构 FoldIndexes: TypeAlias = Union[List[NamedFold], List[TrainTestIndexes]]
步骤2:在原函数中导入并使用别名
import pandas as pd import numpy as np from types import FoldIndexes # 修正原函数语法错误(返回箭头应为`->`) def get_indexes( X: pd.DataFrame, y: pd.Series, ) -> FoldIndexes: """Create indexes to train an estimator. 返回两种可选结构的列表: - 元素为(训练索引数组, 测试索引数组) - 元素为(折名称, (训练索引数组, 测试索引数组)) Args: X (pd.DataFrame): 特征数据集 y (pd.Series): 目标变量 Returns: FoldIndexes: 包含训练测试索引的折列表 """ # 示例实现:生成带名称的5折交叉验证索引 folds = [] total_samples = len(X) for fold_num in range(5): train_idx = np.random.choice(total_samples, int(0.8 * total_samples), replace=False) test_idx = np.setdiff1d(np.arange(total_samples), train_idx) folds.append((f"fold_{fold_num}", (train_idx, test_idx))) return folds
这种方式完全兼容原函数的返回格式,只是通过类型别名把冗长的类型定义简化成FoldIndexes,可读性大幅提升,且types.py可以被任意函数导入复用。
方案2:用标准库数据类实现结构化返回(更直观)
如果想要更贴近你描述的list[fold, [train_indexes, test_index]]结构,可以用Python标准库的dataclasses定义轻量数据类,让返回数据的访问更直观(无需通过元组下标取值)。
步骤1:在types.py中定义数据类
import numpy as np from dataclasses import dataclass from typing import List @dataclass class Fold: """单折交叉验证的索引结构""" name: str | int | None = None # 可选的折名称 train_indexes: np.ndarray # 训练集索引数组 test_indexes: np.ndarray # 测试集索引数组 # 返回类型为Fold对象的列表 FoldList: TypeAlias = List[Fold]
步骤2:修改原函数返回数据类对象
import pandas as pd import numpy as np from types import FoldList, Fold def get_indexes( X: pd.DataFrame, y: pd.Series, ) -> FoldList: """Create indexes to train an estimator. 返回包含Fold对象的列表,每个对象封装了折名称、训练索引和测试索引。 Args: X (pd.DataFrame): 特征数据集 y (pd.Series): 目标变量 Returns: FoldList: 包含折索引信息的对象列表 """ folds = [] total_samples = len(X) for fold_num in range(5): train_idx = np.random.choice(total_samples, int(0.8 * total_samples), replace=False) test_idx = np.setdiff1d(np.arange(total_samples), train_idx) folds.append(Fold(name=f"fold_{fold_num}", train_indexes=train_idx, test_indexes=test_idx)) return folds
兼容原格式的包装器(可选)
如果需要兼容原有代码的返回格式,可以写一个包装器函数转换格式:
from typing import Union, List, Tuple def get_indexes_legacy(X: pd.DataFrame, y: pd.Series) -> Union[List[tuple[str|int, tuple[np.ndarray, np.ndarray]]], List[tuple[np.ndarray, np.ndarray]]]: """兼容原返回格式的包装器""" fold_objects = get_indexes(X, y) legacy_output = [] for fold in fold_objects: if fold.name is not None: legacy_output.append((fold.name, (fold.train_indexes, fold.test_indexes))) else: legacy_output.append((fold.train_indexes, fold.test_indexes)) return legacy_output
这种方式的优势是访问数据更直观,比如fold.train_indexes比fold[1][0]可读性强得多,且完全依赖标准库,无需额外安装依赖。
内容的提问来源于stack exchange,提问作者maxx
相关产品推荐
相关产品推荐

