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

修改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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.19 11:02:53