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

使用rlang结合purrr::pmap编写grouped_lm分组线性回归函数

嘿,这个需求完全可以用rlang的tidy eval工具结合purrr::pmap来实现,我给你写一个完整的grouped_lm函数,还会一步步拆解每个环节的作用,保证你能轻松理解~

实现grouped_lm函数

首先直接上完整函数代码:

library(tidyverse)
library(rlang)

grouped_lm <- function(data, grouping.vars, crit.vars, pred.vars) {
  # 1. 把输入变量转为可用于tidy eval的符号
  grouping_syms <- ensyms(grouping.vars)
  # 提取第一个因变量并转为符号(符合需求规则)
  primary_crit_sym <- ensym(crit.vars[[1]])
  # 把所有自变量转为符号列表
  pred_syms <- ensyms(pred.vars)
  
  # 2. 生成所有分组变量的唯一组合,转成pmap可遍历的格式
  group_combos <- data %>%
    select(!!!grouping_syms) %>%
    distinct() %>%
    pmap(list)  # 每一行分组组合变成键值对列表,比如list(am = 0)
  
  # 3. 遍历每个分组,拟合回归模型
  pmap(group_combos, function(...) {
    # 构建当前分组的筛选条件表达式
    filter_condition <- expr(!!!list(...))
    # 筛选出当前分组的数据子集
    subset_data <- data %>% filter(!!filter_condition)
    
    # 动态构建回归公式:第一个因变量 ~ 所有自变量
    lm_formula <- expr(!!primary_crit_sym ~ !!!pred_syms)
    
    # 拟合线性回归模型
    lm(!!lm_formula, data = subset_data)
  }) %>%
    # 给每个模型命名,用分组的键值对标识(比如"am=0")
    set_names(map_chr(group_combos, ~paste(names(.), ., sep = "=", collapse = "_")))
}

函数功能拆解

  1. 变量符号转换:

    • ensyms()能兼容你输入的裸变量名(比如am)或字符向量(比如"am"),把它们转换成rlang的符号,这样就能在tidy eval环境中灵活引用变量。
    • 严格按照你的需求,只取crit.vars的第一个元素作为回归的因变量。
  2. 生成分组组合:

    • 先提取所有分组变量的唯一值组合,再用pmap(list)把每一行转成一个独立列表,这样后续pmap就能逐个处理每个分组的子集。
  3. 遍历分组拟合模型:

    • 用expr(!!!list(...))把当前分组的键值对转换成筛选表达式(比如am == 0),精准筛选对应的数据子集。
    • 用expr(!!primary_crit_sym ~ !!!pred_syms)动态拼接回归公式,比如当crit.vars首元素是mpg、pred.vars是wt和disp时,会自动生成mpg ~ wt + disp。
    • 最后用!!把公式注入到lm()中完成拟合,再给每个模型加上分组对应的名字,方便后续快速识别。

测试示例(mtcars数据集)

咱们用你提到的场景测试一下函数:

# 调用函数
fit_models <- grouped_lm(
  data = mtcars,
  grouping.vars = am,
  crit.vars = c(mpg, drat),
  pred.vars = c(wt, disp)
)

# 查看自动挡(am=0)的模型结果
fit_models$`am=0`

# 查看手动挡(am=1)的模型结果
fit_models$`am=1`

运行后你会得到两个线性回归模型,分别对应am=0和am=1的分组,每个模型都是用mpg(crit.vars的第一个元素)对wt和disp回归的结果。

扩展:返回整洁的系数结果

如果你想要更便于分析的结构化结果,可以结合broom包修改函数,直接返回包含分组信息的tidy系数表格:

grouped_lm_tidy <- function(data, grouping.vars, crit.vars, pred.vars) {
  grouping_syms <- ensyms(grouping.vars)
  primary_crit_sym <- ensym(crit.vars[[1]])
  pred_syms <- ensyms(pred.vars)
  
  group_combos <- data %>%
    select(!!!grouping_syms) %>%
    distinct() %>%
    pmap(list)
  
  pmap_dfr(group_combos, function(...) {
    filter_condition <- expr(!!!list(...))
    subset_data <- data %>% filter(!!filter_condition)
    
    lm_formula <- expr(!!primary_crit_sym ~ !!!pred_syms)
    
    lm(!!lm_formula, data = subset_data) %>%
      broom::tidy() %>%
      mutate(!!!list(...))  # 把分组变量添加到结果表格中
  })
}

# 调用扩展函数
tidy_results <- grouped_lm_tidy(
  data = mtcars,
  grouping.vars = am,
  crit.vars = c(mpg, drat),
  pred.vars = c(wt, disp)
)

print(tidy_results)

这样返回的是一个数据框,每个分组的系数、标准误、p值等信息都清晰列出来了,非常适合后续的可视化或统计分析。

内容的提问来源于stack exchange,提问作者Indrajeet Patil

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.21 03:32:48