如何通过循环计算不同剪枝决策树的训练与测试误差?
决策树剪枝后计算训练/测试误差的问题解决
问题说明
需要对决策树进行剪枝,生成终端节点数为2-20的19棵树,并计算每棵树的训练和测试均方误差(MSE)。已成功生成剪枝树,但添加误差计算逻辑后代码报错。
你尝试的错误代码示例
仅生成剪枝树的可行代码
range <- c(2:20) for (i in range) { prune.fit <- prune.tree(fit, best = i) plot(prune.fit) # 绘制每棵树 text(prune.fit, pretty = 0) }
添加误差计算后的错误尝试
# 错误尝试1 for (i in range) { pred.fittrain[i] <- predict(prune.fit[i], newdata = my_ahp_train) mean((pred.fittrain - my_ahp_train$sale_price)^2) pred.fittest[i] <- predict(prune.fit[i], newdata = my_ahp_test) mean((pred.fittest - my_ahp_test$sale_price)^2) } # 错误尝试2 range <- c(2:20) for (i in range) { prune.fit <- prune.tree(fit, best = i) plot(prune.fit) text(prune.fit, pretty = 0) pred.fittrain[i] <- predict(prune.fit[i], newdata = my_ahp_train) mean((pred.fittrain - my_ahp_train$sale_price)^2) pred.fittest[i] <- predict(prune.fit[i], newdata = my_ahp_test) mean((pred.fittest - my_ahp_test$sale_price)^2) }
错误原因分析
- 模型对象索引错误:
prune.fit是单次循环生成的单个剪枝树模型,不是列表,不能用prune.fit[i]的方式访问,直接使用prune.fit即可。 - 未存储误差结果:仅计算MSE但未将结果保存到变量中,无法后续查看每棵树的误差。
- 预测结果存储逻辑错误:
pred.fittrain[i]试图将整个预测向量存入单个元素,会导致维度不匹配的错误;我们需要的是误差值,无需保存完整预测结果(若需保存,可改用列表)。
修正后的完整代码
# 定义终端节点数量范围 range <- 2:20 # 初始化存储训练、测试MSE的向量 train_mse <- numeric(length(range)) test_mse <- numeric(length(range)) # 循环生成剪枝树并计算误差 for (idx in seq_along(range)) { node_count <- range[idx] # 生成指定终端节点数的剪枝树 prune.fit <- prune.tree(fit, best = node_count) # 绘制剪枝树(可选步骤) plot(prune.fit) text(prune.fit, pretty = 0) # 计算训练集MSE pred_train <- predict(prune.fit, newdata = my_ahp_train) train_mse[idx] <- mean((pred_train - my_ahp_train$sale_price)^2) # 计算测试集MSE pred_test <- predict(prune.fit, newdata = my_ahp_test) test_mse[idx] <- mean((pred_test - my_ahp_test$sale_price)^2) } # 整理并输出结果 error_results <- data.frame( 终端节点数 = range, 训练集MSE = train_mse, 测试集MSE = test_mse ) print(error_results) # 可选:可视化误差随终端节点数的变化 plot(range, train_mse, type = "l", col = "blue", xlab = "终端节点数", ylab = "均方误差(MSE)", main = "剪枝树的训练与测试误差对比") lines(range, test_mse, col = "red") legend("topright", legend = c("训练MSE", "测试MSE"), col = c("blue", "red"), lty = 1)
代码说明
- 提前初始化
train_mse和test_mse向量,用于存储每棵树的误差值。 - 循环中用
seq_along(range)获取索引,避免终端节点数与向量索引混淆。 - 直接使用
prune.fit调用predict函数,修正模型对象的访问方式。 - 计算MSE后直接存入对应向量位置,确保结果被保存。
- 最后将结果整理为数据框,方便查看和后续分析,也可选择可视化误差变化趋势。
内容的提问来源于stack exchange,提问作者madibecoding
相关产品推荐
相关产品推荐

