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

如何手动复现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")的计算逻辑是:

  1. 按nodegree的取值将原始数据分组;
  2. 在每个分组内,为treat的所有可能水平(0和1)分别生成新数据集——保留该分组内其他变量(如married)的原始观测值;
  3. 对每个新数据集计算所有观测的预测值,再取平均值,最终输出每个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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 16:57:30