R语言决策树剪枝----泰坦尼克数据集

R 复制代码
# 泰坦尼克号 决策树案例 (树层级多)

# 安装包
# install.packages(c("rpart", "rpart.plot", "titanic"))

library(rpart)       # 决策树
library(rpart.plot)   # 画图
library(titanic)      # 泰坦尼克数据集

# 加载数据
data("titanic_train")
df <- titanic_train
head(df)
# 数据预处理(决策树必须处理缺失值)
df$Age[is.na(df$Age)] <- median(df$Age, na.rm=TRUE)
df$Embarked[df$Embarked==""] <- "S"

# 把分类变量转成因子
df$Survived <- as.factor(df$Survived)
df$Pclass <- as.factor(df$Pclass)
df$Sex <- as.factor(df$Sex)
df$Embarked <- as.factor(df$Embarked)

训练模型

计算存活情况

R 复制代码
# --------------------
# 训练 深层决策树(层级多)
# --------------------
tree_full <- rpart(
  Survived ~ Pclass + Sex + Age + SibSp + Parch + Fare + Embarked,
  data = df,
  method = "class",
  cp = 0,          # cp=0 → 不剪枝,树长到最复杂
  maxdepth = 8     # 允许树长得很深
)

# 画图:层级非常多!
rpart.plot(
  tree_full,
  main = "泰坦尼克号 深层决策树(层级多)",
  type = 4,
  extra = 101,
  cex = 0.7        # 字体缩小,才能显示多层
)

剪枝

CP = Complexity Parameter 复杂度参数

  • CP 越大 → 树越简单
  • CP 越小 → 树越复杂
  • CP=0 → 不剪枝,树长到最大
R 复制代码
print(tree_deep$cptable)
#找到 xerror 最小的那一行,取它的 CP
print(min(tree_deep$cptable[,'xerror']))

> print(tree_full$cptable)

CP nsplit rel error xerror xstd

1 0.444444444 0 1.0000000 1.0000000 0.04244576

2 0.030701754 1 0.5555556 0.5555556 0.03574957

3 0.023391813 3 0.4941520 0.5029240 0.03444798

4 0.020467836 4 0.4707602 0.4912281 0.03413963

5 0.014619883 5 0.4502924 0.4853801 0.03398272

6 0.007309942 6 0.4356725 0.5058480 0.03452394

7 0.006578947 10 0.4035088 0.4912281 0.03413963

8 0.004385965 14 0.3771930 0.4766082 0.03374384

9 0.002923977 16 0.3684211 0.4649123 0.03341867

10 0.000000000 18 0.3625731 0.4707602 0.03358222

R 复制代码
best_cp <- tree_full$cptable[which.min(tree_full$cptable[, "xerror"]), "CP"]
tree_pruned <- prune(tree_full, cp = 0.03)

rpart.plot(
  tree_pruned,
  main = "泰坦尼克号 深层决策树",
  type = 4,
  extra = 101,
  cex = 0.7  
)
相关推荐
临床数据科学和人工智能兴趣组6 小时前
在R语言中, 使用 as.factor() 函数转换数值型变量
开发语言·r语言
为啥全要学1 天前
PyTorch 中网络剪枝、梯度剪裁、梯度累积
网络·pytorch·剪枝
babe小鑫10 天前
信息与计算科学专业应届生面试 怎么证明自己能解决业务问题
学习·r语言·excel
木井巳10 天前
【记忆化搜索】最长递增子序列
java·算法·leetcode·深度优先·剪枝·推荐算法
毕业设计70310 天前
(免费领源码) SpringBoot 游戏交易平台17600-java、PHP、python、C#、小程序、大数据、单片机、网络工程等)
java·spring boot·mysql·决策树·mybatis·idea·推荐算法
木井巳11 天前
【记忆化搜索】不同路径
java·算法·leetcode·深度优先·剪枝·推荐算法
深兰科技11 天前
深兰科技受邀参与第二届中国(南宁)—东盟人工智能场景应用对接会,深化AI国际合作
人工智能·qt·r语言·scala·symfony·智能机器人·深兰科技
赵钰老师12 天前
基于ArcGIS Pro、R、INVEST等多技术融合下生态系统服务权衡与协同动态分析
python·arcgis·数据分析·r语言
User_芊芊君子16 天前
RStudio 鸿蒙 PC 适配全记录:以 Qt 原生工作区承载嵌入式 R
qt·r语言·harmonyos
临床数据科学和人工智能兴趣组16 天前
399元现在超值!学R语言,订阅我们专栏就够了,包括了所有的内容,不断更新!
人工智能·数据挖掘·r语言·r语言-4.2.1·临床研究