PySpark或MLlib中是否有对应scikit-learn的train_test_split功能?
PySpark 等效 sklearn train_test_split 的实现方案
PySpark 原生提供了和 sklearn.model_selection.train_test_split 功能完全对齐的 API:DataFrame 自带的 randomSplit() 方法,支持指定拆分占比、随机种子保证结果可复现,完全可以满足你的需求。
前置准备
PySpark 机器学习流程中通常将特征列和标签列放在同一个 DataFrame 中处理,所以你需要先把单独的特征数据集 featuresonly 和标签数据集 target 合并为一个总 DataFrame,如果你的数据已经是合并状态可以直接跳过这一步:
from pyspark.sql import SparkSession from pyspark.sql.functions import monotonically_increasing_id # 初始化SparkSession,已初始化的话可以省略 spark = SparkSession.builder.appName("dataset_split").getOrCreate() # 给两个数据集添加行索引后合并,保证特征和标签一一对应 features_df = featuresonly.withColumn("row_id", monotonically_increasing_id()) target_df = target.withColumn("row_id", monotonically_increasing_id()) # 合并后默认特征列名为features,标签列名为label,可根据实际情况修改字段名 total_df = features_df.join(target_df, on="row_id", how="inner").drop("row_id")
数据集拆分完整实现
完全对应你给出的 sklearn 代码逻辑:测试集占比20%、随机种子123,PySpark 实现代码如下:
# 拆分数据集,参数依次为[训练集占比, 测试集占比]、随机种子(和sklearn的random_state作用一致) train_df, test_df = total_df.randomSplit([0.8, 0.2], seed=123) # 统计拆分结果,和原打印逻辑完全对齐 train_total = train_df.count() test_total = test_df.count() # 原代码len(y_train)-y_train.sum()对应二分类场景下标签为0的样本数量,可根据实际标签定义调整过滤条件 train_good = train_df.filter(train_df.label == 0).count() test_good = test_df.filter(test_df.label == 0).count() print("Training set has {} samples.".format(train_total)) print("Testing set has {} samples.".format(test_total)) print("Training set has good {} samples.".format(train_good)) print("Testing set has good {} samples.".format(test_good))
补充说明
- 如果你需要实现分层抽样(和 sklearn 的
stratify参数功能对齐),可以使用sampleBy方法按标签占比拆分,避免拆分后数据集分布偏移。 randomSplit返回的 DataFrame 是惰性计算的,只有调用count()这类行动算子时才会真正执行拆分逻辑。
内容的提问来源于stack exchange,提问作者Nabih Bawazir
相关产品推荐
相关产品推荐

