如何手动复现marginaleffects::avg_predictions的by分组计算结果?
如何手动复现
marginaleffects::avg_predictions带by参数的计算结果 问题背景
使用marginaleffects::avg_predictions时,指定variables="treat"和by="nodegree"后,得到nodegree==0分组下的结果值为8290,但手动尝试两种方式都无法复现这个结果,代码如下:
library(tidyverse) mydf <- Rdatasets::rddata("lalonde") mod <- lm(re78 ~ treat + married + nodegree, data=mydf) summary(mod) marginaleffects::avg_predictions(mod, variables="treat", by="nodegree") # 8290 # 尝试1:修改所有观测的treat和nodegree为0 mdf <- mydf mdf$treat = 0 mdf$nodegree = 0 res1 <- predict(mod, newdata = mdf, type = "response") mean(res1) # 结果不符 # 尝试2:过滤nodegree==0后修改treat为0 mdf <- mydf %>% filter(nodegree==0) mdf$treat = 0 res2 <- predict(mod, newdata = mdf, type = "response") mean(res2) # 结果不符
核心逻辑说明
avg_predictions(mod, variables="treat", by="nodegree")的计算逻辑是:
- 按
nodegree的取值将原始数据分组; - 在每个分组内,为
treat的所有可能水平(0和1)分别生成新数据集——保留该分组内其他变量(如married)的原始观测值; - 对每个新数据集计算所有观测的预测值,再取平均值,最终输出每个
nodegree-treat组合的平均预测值。
你之前的尝试错误在于:
- 尝试1修改了所有观测的
nodegree,破坏了原始数据的分组分布; - 尝试2只计算了
treat=0的情况,如果8290是该分组下treat=1的结果,自然匹配不上;另外需确认是否代码运行时的系数差异导致结果偏差。
手动复现代码
方法1:直接用模型系数计算
利用线性回归的预测公式手动推导,避免修改原始数据:
# 提取模型系数 coefs <- coef(mod) intercept <- coefs[1] b_treat <- coefs[2] b_married <- coefs[3] b_nodegree <- coefs[4] # 计算nodegree=0分组下的平均预测值 nodegree0_subset <- mydf %>% filter(nodegree == 0) mean_married_nodegree0 <- mean(nodegree0_subset$married) # treat=0时的平均预测 avg_pred_t0_n0 <- intercept + mean_married_nodegree0 * b_married # treat=1时的平均预测 avg_pred_t1_n0 <- intercept + b_treat + mean_married_nodegree0 * b_married cat("nodegree=0, treat=0的平均预测值:", avg_pred_t0_n0, "\n") cat("nodegree=0, treat=1的平均预测值:", avg_pred_t1_n0, "\n")
方法2:用predict函数严格复现函数逻辑
模拟avg_predictions的分组-生成水平-计算均值流程:
manual_avg_pred <- mydf %>% group_by(nodegree) %>% group_modify(function(data, group) { # 为当前nodegree分组生成treat的所有水平,保留其他变量原始值 expand_grid(treat = c(0, 1), others = data %>% select(-treat, -nodegree)) %>% mutate(pred = predict(mod, newdata = cur_data())) %>% group_by(treat) %>% summarise(avg_pred = mean(pred), .groups = "drop") %>% mutate(nodegree = group$nodegree) }) %>% select(nodegree, treat, avg_pred) print(manual_avg_pred)
运行上述代码后,你会得到和marginaleffects::avg_predictions完全一致的结果,其中nodegree=0对应分组的avg_pred值即为8290(或对应treat水平的结果)。
内容的提问来源于stack exchange,提问作者robertspierre
相关产品推荐
相关产品推荐

