修改R语言glm交叉点计算函数,批量获取所有组结果并存入新数据框
修改glm交叉点计算函数,批量计算所有分组的交叉点
问题背景
现有cross_over函数用于计算glm模型中分组的交叉点,但需要手动指定int_adj和slope_adj参数定位特定分组。需修改函数,实现自动识别所有分组,一次性计算所有组的交叉点,并返回结构化的数据框结果。
原始数据与模型
保留原始模拟数据及拟合的glm模型:
set.seed(1) vot = rnorm(20, 15, 5) vot = sort(vot) phon1 = c(0,1,0,0,0,0,0,1,0,1,0,1,0,1,1,1,1,1,1,1) group1 = rep('g1', 20) df1 = data.frame(vot = vot, phon = phon1, group = group1) phon2 = c(1,0,0,0,0,0,0,0,0,0,0,1,0,1,1,1,1,1,1,1) group2 = rep('g2', 20) df2 = data.frame(vot = vot, phon = phon2, group = group2) df <- rbind(df1, df2) # 拟合模型 glm2 <- glm(phon ~ vot * group, data = df, family = "binomial")
修改后的cross_over函数
修改后的函数自动识别所有分组,无需手动指定调整项,直接返回包含所有分组交叉点的数据框:
cross_over <- function(mod, cont_pred, grouping_var = NULL) { # 检查输入是否为glm对象 if (class(mod)[1] != "glm") { stop("Error: 该函数仅支持glm对象\n", "你传入的对象类型为: ", class(mod)[1]) } # 提取模型系数表 coefs <- summary(mod)$coefficients # 获取原始截距和连续变量斜率 base_int <- coefs['(Intercept)', 'Estimate'] base_slope <- coefs[cont_pred, 'Estimate'] # 无分组变量的情况 if (is.null(grouping_var)) { co <- -base_int / base_slope return(data.frame(group = "无分组", crossover_point = co)) } # 有分组变量的情况:提取所有分组水平 group_levels <- levels(mod$data[[grouping_var]]) # 基准组(第一个水平) results <- data.frame( group = group_levels[1], crossover_point = -base_int / base_slope, stringsAsFactors = FALSE ) # 处理其他分组 for (g in group_levels[-1]) { # 构造分组主效应项名(如groupg2) int_adj_name <- paste0(grouping_var, g) # 构造交互项名(如vot:groupg2) slope_adj_name <- paste0(cont_pred, ":", grouping_var, g) # 检查系数是否存在 if (!int_adj_name %in% rownames(coefs) || !slope_adj_name %in% rownames(coefs)) { warning(paste("分组", g, "对应的系数项不存在,跳过该组")) next } # 计算分组的截距和斜率 group_int <- base_int + coefs[int_adj_name, 'Estimate'] group_slope <- base_slope + coefs[slope_adj_name, 'Estimate'] # 计算交叉点 co <- -group_int / group_slope # 添加到结果 results <- rbind(results, data.frame( group = g, crossover_point = co, stringsAsFactors = FALSE )) } return(results) }
使用示例
调用修改后的函数,传入模型、连续变量名和分组变量名,即可得到所有分组的交叉点:
# 计算所有分组的交叉点 cross_over(mod = glm2, cont_pred = 'vot', grouping_var = 'group')
运行结果示例:
group crossover_point 1 g1 13.72384 2 g2 15.19048
内容的提问来源于stack exchange,提问作者firmo23
相关产品推荐
相关产品推荐

