如何使用R tidymodels复现plot.lda()函数的可视化效果
优化实现方案
这里提供两种更简洁稳定的实现思路,避免手动硬编码提取预测结果的问题:
方案1:复用broom对LDA对象的原生支持,步骤最少
broom包已经适配了MASS::lda的拟合结果,可以直接用augment()生成绑定了预测值的数据集,不需要手动调用predict()做拼接:
library(broom) # 直接处理原生lda拟合对象,自动生成包含判别值、分类结果的数据集 lda_aug <- augment(lda.fit, data = Smarket_train) ggplot(lda_aug, aes(x = .fitted)) + # 若要完全对齐base R plot.lda的默认密度曲线效果,用geom_density即可 geom_density(fill = "gray85", color = "black", linewidth = 0.8) + # 保留你的直方图写法的话替换为下行即可,注意用after_stat适配新版ggplot2 # geom_histogram(aes(y = after_stat(density)), binwidth = 0.5, alpha = 0.7, fill = "steelblue") + scale_x_continuous(breaks = seq(-4, 4, by = 2)) + facet_grid(vars(Direction)) + labs(x = "线性判别得分", y = "密度")
方案2:全tidymodels规范实现,无硬编码
如果要完全遵循tidymodels的工作流逻辑,避免直接操作底层拟合对象,可以通过按列名提取判别值代替按索引[[3]]的硬编码写法,后续模型调整时兼容性更强:
# 直接用workflow做预测,不需要提前提取拟合对象 pred_lda <- predict(the_workflow, new_data = Smarket_train, type = "raw") %>% as_tibble() %>% pull(x) # 按列名提取判别值,比按索引取更稳定 Smarket_train %>% bind_cols(.fitted = pred_lda) %>% ggplot(aes(x = .fitted)) + geom_histogram(aes(y = after_stat(density)), binwidth = 0.5, alpha = 0.7) + scale_x_continuous(breaks = seq(-4, 4, by = 2)) + facet_grid(vars(Direction)) + labs(x = "", y = "密度")
关键优化点
- 去掉了按索引提取预测结果的硬编码逻辑,后续如果模型输入特征、判别维度发生变化,代码不会失效
- 用
after_stat(density)替代了已经废弃的stat(density)写法,适配ggplot2 3.3.0之后的所有版本 - 方案1完全省去了手动处理预测结果、拼接数据集的步骤,代码更简洁易读
内容的提问来源于stack exchange,提问作者itsMeInMiami
相关产品推荐
相关产品推荐

