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

brms模型边际效应链诊断:如何提取链计算Rhat值?

分state的贝叶斯边际处理效应链诊断方法

问题背景

用brms拟合带state层面随机斜率的logit模型:

fit <- brm(y ~ (treat|state) + treat + age + sex + race, family = "bernoulli", data = dat)

通过marginaleffects::avg_slopes计算分state的边际处理效应(MTE):

mfx <- avg_slopes(fit, variable = "treat", by = "state")

需要对每个state的MTE做链诊断(如计算Rhat值),但无法直接用posterior_draws提取对应链的样本。


解决方案:提取样本层面的MTE并计算诊断指标

方法1:用marginaleffects::comparisons获取带链信息的原始样本

comparisons函数可直接输出每个后验样本、每个state的处理效应,结合brms模型的链标识信息后即可计算Rhat:

# 1. 获取所有后验样本的分state处理效应
comp <- comparisons(
    fit,
    variable = "treat",
    by = "state",
    draws = "all",  # 保留所有后验样本
    type = "response"  # 若需要链接尺度效应,替换为"link"
)

# 2. 从brms模型中提取每个样本对应的链信息
chain_info <- as.data.frame(fit, variable = "lp__") %>% 
    mutate(chain = as.integer(stringr::str_extract(.draw, "chain:(\\d+)") %>% stringr::str_remove("chain:")))

# 3. 合并处理效应样本与链信息
comp_with_chain <- comp %>% 
    left_join(chain_info, by = ".draw")

# 4. 按state分组计算Rhat
library(posterior)
rhat_results <- comp_with_chain %>% 
    group_by(state) %>% 
    summarise(
        rhat = rhat(matrix(estimate, ncol = length(unique(chain)))),
        .groups = "drop"
    )

方法2:从brms后验样本手动计算MTE

模型中每个state的处理效应为treat的固定效应 + 对应state的treat随机斜率,可直接提取参数计算:

# 1. 提取固定效应和随机效应的后验样本
post <- posterior_samples(fit)

# 2. 提取所有state的treat随机斜率列(列名格式为r_state__treat[state名称])
rand_slope_cols <- grep("^r_state__treat", colnames(post), value = TRUE)
state_names <- stringr::str_remove(rand_slope_cols, "^r_state__treat\\[|\\]$")

# 3. 计算每个state的MTE样本:固定效应 + 对应随机斜率
mte_samples <- data.frame(
    .draw = 1:nrow(post),
    chain = rep(1:fit$fit@sim$chains, each = fit$fit@sim$iter - fit$fit@sim$warmup)
)
for (i in seq_along(state_names)) {
    mte_samples[[state_names[i]]] <- post$b_treat + post[[rand_slope_cols[i]]]
}

# 4. 按state分组计算Rhat
rhat_results <- mte_samples %>% 
    select(-.draw, -chain) %>% 
    pivot_longer(everything(), names_to = "state", values_to = "estimate") %>% 
    group_by(state) %>% 
    summarise(
        rhat = rhat(matrix(estimate, ncol = length(unique(mte_samples$chain)))),
        .groups = "drop"
    )

内容的提问来源于stack exchange,提问作者Jonathan Gellar

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 22:07:44