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

请求协助:创建基于元组参数的训练测试集拆分函数

数据集拆分函数实现

你可以使用scikit-learn中的train_test_split工具来实现这个需求,以下是完整的函数代码:

from sklearn.model_selection import train_test_split
import pandas as pd

def split_dataset(data_tuple):
    X, y = data_tuple
    # 可选:校验输入类型,提前排查不符合要求的输入
    if not isinstance(X, pd.DataFrame):
        raise TypeError("X必须是pandas DataFrame类型")
    if not isinstance(y, pd.Series):
        raise TypeError("y必须是pandas Series类型")
    
    # 按70%训练集、30%验证集拆分,random_state保证拆分结果可复现
    X_train, X_test, y_train, y_test = train_test_split(
        X, y, test_size=0.3, random_state=42
    )
    return (X_train, X_test, y_train, y_test)

关键说明:

  • train_test_split会自动保留输入的原始类型,因此只要传入的X是DataFrame、y是Series,返回的结果也会对应保持该类型
  • 设置random_state=42是为了让每次拆分的结果一致,方便调试和复现,你可以根据需要修改这个值或者直接移除该参数
  • 可选的类型校验可以提前发现输入不符合要求的情况,避免后续流程报错

使用示例:

# 构造示例数据集
X = pd.DataFrame({'feature1': [1,2,3,4,5,6,7,8,9,10], 'feature2': [11,12,13,14,15,16,17,18,19,20]})
y = pd.Series([0,1,0,1,0,1,0,1,0,1])

# 调用函数拆分数据集
X_train, X_test, y_train, y_test = split_dataset((X, y))

# 验证结果类型
print(type(X_train))  # <class 'pandas.core.frame.DataFrame'>
print(type(y_train))  # <class 'pandas.core.series.Series'>

内容的提问来源于stack exchange,提问作者tt241

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 22:25:25