如何基于分组变量从rstanarm包的stan_glm()获取后验预测?
针对rstanarm中分组变量的后验预测方法
嗨,我来帮你搞定这个问题!要从stan_glm()的结果里针对vs的两个分组(0和1)分别获取后验预测,核心在于构造合适的newdata参数——这个数据集需要清晰定义你想要预测的每个分组的特征值,下面分两种常见场景给你具体代码:
场景1:获取每组的平均水平后验预测
如果你想得到vs=0和vs=1两组各自的平均预测结果,我们需要先构造包含每组自变量均值(或典型值)的数据集:
方法1:用dplyr分组汇总(推荐,代码更简洁)
library(rstanarm) library(dplyr) # 先拟合你的模型(和你写的一致) fit <- stan_glm(mpg ~., data = mtcars) # 按vs分组,计算所有自变量的均值(分类变量如am也可以用众数,这里用均值简化) grouped_newdata <- mtcars %>% group_by(vs) %>% summarise(across(everything(), mean)) # 针对分组数据做后验预测 post_pred_groups <- posterior_predict(fit, newdata = grouped_newdata)
方法2:用基础R实现(无需额外包)
library(rstanarm) fit <- stan_glm(mpg ~., data = mtcars) # 按vs分组,计算各列均值 grouped_newdata <- aggregate(. ~ vs, data = mtcars, FUN = mean) post_pred_groups <- posterior_predict(fit, newdata = grouped_newdata)
解读结果
post_pred_groups是一个矩阵:每一行代表一个后验样本,每一列对应一个分组(第一列是vs=0,第二列是vs=1)。你可以这样分析每组的预测分布:
# 分析vs=0组的后验预测 vs0_pred <- post_pred_groups[, 1] cat("vs=0组的预测均值:", mean(vs0_pred), "\n") cat("vs=0组的95%可信区间:", quantile(vs0_pred, c(0.025, 0.975)), "\n") # 分析vs=1组的后验预测 vs1_pred <- post_pred_groups[, 2] cat("vs=1组的预测均值:", mean(vs1_pred), "\n") cat("vs=1组的95%可信区间:", quantile(vs1_pred, c(0.025, 0.975)), "\n")
场景2:获取每个原始观测的后验预测并按分组拆分
如果你想对mtcars里的每个观测都做预测,再按vs分组整理结果,可以直接传入原始数据集:
library(rstanarm) fit <- stan_glm(mpg ~., data = mtcars) # 对所有原始观测做后验预测 post_pred_all <- posterior_predict(fit, newdata = mtcars) # 按vs分组提取预测结果 vs0_pred_individual <- post_pred_all[, mtcars$vs == 0] # vs=0的所有观测的预测 vs1_pred_individual <- post_pred_all[, mtcars$vs == 1] # vs=1的所有观测的预测
这样你就可以针对每组的个体预测做进一步分析,比如计算组内每个观测的预测分布统计量。
内容的提问来源于stack exchange,提问作者rnorouzian
相关产品推荐
相关产品推荐

