使用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 = "_"))) }
函数功能拆解
变量符号转换:
ensyms()能兼容你输入的裸变量名(比如am)或字符向量(比如"am"),把它们转换成rlang的符号,这样就能在tidy eval环境中灵活引用变量。- 严格按照你的需求,只取
crit.vars的第一个元素作为回归的因变量。
生成分组组合:
- 先提取所有分组变量的唯一值组合,再用
pmap(list)把每一行转成一个独立列表,这样后续pmap就能逐个处理每个分组的子集。
- 先提取所有分组变量的唯一值组合,再用
遍历分组拟合模型:
- 用
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
相关产品推荐
相关产品推荐

