【问题标题】:Cross-validating a CART model交叉验证 CART 模型
【发布时间】:2013-05-19 01:14:15
【问题描述】:

在一项作业中,我们被要求对 CART 模型执行交叉验证。我曾尝试使用cvTools 中的cvFit 函数,但收到一条奇怪的错误消息。这是一个最小的例子:

library(rpart)
library(cvTools)
data(iris)
cvFit(rpart(formula=Species~., data=iris))

我看到的错误是:

Error in nobs(y) : argument "y" is missing, with no default

还有traceback()

5: nobs(y)
4: cvFit.call(call, data = data, x = x, y = y, cost = cost, K = K, 
       R = R, foldType = foldType, folds = folds, names = names, 
       predictArgs = predictArgs, costArgs = costArgs, envir = envir, 
       seed = seed)
3: cvFit(call, data = data, x = x, y = y, cost = cost, K = K, R = R, 
       foldType = foldType, folds = folds, names = names, predictArgs = predictArgs, 
       costArgs = costArgs, envir = envir, seed = seed)
2: cvFit.default(rpart(formula = Species ~ ., data = iris))
1: cvFit(rpart(formula = Species ~ ., data = iris))

看起来ycvFit.default 的必填项。但是:

> cvFit(rpart(formula=Species~., data=iris), y=iris$Species)
Error in cvFit.call(call, data = data, x = x, y = y, cost = cost, K = K,  : 
  'x' must have 0 observations

我做错了什么?哪个包可以让我对 CART 树进行交叉验证,而无需自己编写代码? (我太懒了……)

【问题讨论】:

  • 如果您深入了解 cvTools 的文档,似乎大多数这些工具都是在构建时考虑到连续响应变量,而不是离散的。你或许可以让它工作,但看起来你必须向cost 提供你自己的函数来计算分类错误。
  • @joran:没错——谢谢!见my own answer

标签: r cross-validation rpart


【解决方案1】:

插入符号包使交叉验证变得轻而易举:

> library(caret)
> data(iris)
> tc <- trainControl("cv",10)
> rpart.grid <- expand.grid(.cp=0.2)
> 
> (train.rpart <- train(Species ~., data=iris, method="rpart",trControl=tc,tuneGrid=rpart.grid))
150 samples
  4 predictors
  3 classes: 'setosa', 'versicolor', 'virginica' 

No pre-processing
Resampling: Cross-Validation (10 fold) 

Summary of sample sizes: 135, 135, 135, 135, 135, 135, ... 

Resampling results

  Accuracy  Kappa  Accuracy SD  Kappa SD
  0.94      0.91   0.0798       0.12    

Tuning parameter 'cp' was held constant at a value of 0.2

【讨论】:

  • 哇。只需查看train 中支持的方法列表即可。这就是我所说的全面......这里发生了很多“魔法”。是否可以只访问交叉验证例程,而不实际优化模型参数?
  • 我不这么认为,但您可以定义自己的参数网格。如果您不想测试多个模型,则可以将它们设置为静态值。我将通过编辑上面的示例来说明这一点。
  • 什么是插入符号?我没有看到您的答案中使用了它。
  • 一个我忘记包含在代码中的库,进行了编辑,所以现在应该全部设置好了。
【解决方案2】:

最后,我能够让它工作。正如 Joran 所指出的,需要调整 cost 参数。在我的情况下,我使用的是 0/1 损失,这意味着我使用一个简单的函数来评估 != 而不是 yyHat 之间的 -。此外,predictArgs 必须包含c(type='class'),否则内部使用的predict 调用将返回概率向量而不是最可能的分类。总结一下:

library(rpart)
library(cvTools)
data(iris)
cvFit(rpart, formula=Species~., data=iris,
      cost=function(y, yHat) (y != yHat) + 0, predictArgs=c(type='class'))

(这使用了cvFit 的另一种变体。可以通过设置args= 参数来传递rpart 的其他参数。)

【讨论】:

    猜你喜欢
    • 2014-02-18
    • 2013-12-08
    • 2016-05-26
    • 2020-03-26
    • 2015-12-22
    • 2015-05-01
    • 2017-01-27
    • 2021-06-09
    • 1970-01-01
    相关资源
    最近更新 更多