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

如何通过循环计算不同剪枝决策树的训练与测试误差?

决策树剪枝后计算训练/测试误差的问题解决

问题说明

需要对决策树进行剪枝,生成终端节点数为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)
}

错误原因分析

  1. 模型对象索引错误:prune.fit是单次循环生成的单个剪枝树模型,不是列表,不能用prune.fit[i]的方式访问,直接使用prune.fit即可。
  2. 未存储误差结果:仅计算MSE但未将结果保存到变量中,无法后续查看每棵树的误差。
  3. 预测结果存储逻辑错误: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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 22:52:50