请求协助:创建基于元组参数的训练测试集拆分函数
数据集拆分函数实现
你可以使用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
相关产品推荐
相关产品推荐

