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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 15:24:03