You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何简化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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.05 15:27:11