caret包createDataPartition函数中y参数的作用及指定原因?
关于caret包中
createDataPartition()的y参数作用 createDataPartition()的核心是做分层采样,而y参数就是用来定义分层依据的——哪怕它返回的是行索引,这些索引的选择逻辑完全依赖y的分布。
为什么需要分层采样?
简单随机采样可能会导致采样偏差:比如分类任务里某类样本在训练集里占比过低,回归任务里训练集的y值范围和原数据偏差太大,这都会让模型泛化能力变差。y参数就是让函数按照y的分布来选行,保证采样后的子集和原数据的y分布一致。
具体场景示例
1. 分类任务(y为类别变量)
假设你用鸢尾花数据集,y是Species(三类):
library(caret) data(iris) # 按Species分层采样,取70%作为训练集 train_idx <- createDataPartition(y = iris$Species, p = 0.7, list = FALSE) train_data <- iris[train_idx, ] test_data <- iris[-train_idx, ] # 查看两类数据集的类别比例 table(train_data$Species) #> setosa versicolor virginica #> 35 35 35 table(test_data$Species) #> setosa versicolor virginica #> 15 15 15
可以看到训练集和测试集的三类样本比例和原数据完全一致(原数据每类50个)。如果用简单随机采样,大概率会出现某类样本数量偏离的情况。
2. 回归任务(y为连续变量)
对于回归任务,函数会自动把y分成若干分位数组,再按组比例采样,保证子集的y分布和原数据接近:
data(mtcars) # 按mpg(连续变量)分层采样 train_idx_reg <- createDataPartition(y = mtcars$mpg, p = 0.7, list = FALSE) train_reg <- mtcars[train_idx_reg, ] test_reg <- mtcars[-train_idx_reg, ] # 对比mpg的统计特征 summary(train_reg$mpg) #> Min. 1st Qu. Median Mean 3rd Qu. Max. #> 10.4 15.2 19.2 20.1 22.8 33.9 summary(test_reg$mpg) #> Min. 1st Qu. Median Mean 3rd Qu. Max. #> 11.0 15.5 18.1 19.7 21.4 32.4
可以看到训练集和测试集的mpg分布特征(均值、分位数)非常接近,远好于简单随机采样的结果。
总结
y参数的本质是告诉函数:按照这个变量的分布来划分数据集,确保采样后的子集保留原数据的关键分布特征,避免采样偏差。函数返回的行索引只是最终结果,背后的采样逻辑完全依赖y的信息。
内容的提问来源于stack exchange,提问作者Preston Lui
相关产品推荐
相关产品推荐

