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

如何对DALEX中多模型的model_profile连续变量结果求平均?

处理DALEX中多个continuous类型model_profile的均值计算

针对你提到的500个模型的连续变量部分依赖图(partial dependence profile)求平均需求,分两种场景给出解决方案:

场景1:尚未生成所有model_profile(推荐)

通过统一的x网格生成所有模型的profile,确保变量的x值完全对齐,后续直接分组求平均即可,准确性更高。

实现步骤:

  1. 基于完整数据集生成统一的x网格:连续变量取分位点,分类变量保留所有类别
  2. 生成每个模型的model_profile时指定该网格
  3. 合并所有profile数据,按变量、x值分组计算平均预测值
library("ranger") 
library(DALEX)
library(dplyr)

# 加载内置数据
data(titanic_imputed)

# 生成统一的x网格:连续变量取100个分位点,分类变量取所有唯一值
common_grid <- lapply(titanic_imputed[, setdiff(names(titanic_imputed), "survived")], function(col) {
  if (is.numeric(col)) {
    quantile(col, seq(0, 1, length.out = 100))
  } else {
    unique(col)
  }
})

# ----------------------
# 示例:生成2个模型的profile(实际替换为500个模型的循环/批量处理)
# ----------------------
# 模型1
trainIdx_1 <- sample(nrow(titanic_imputed), 2/3 * nrow(titanic_imputed)) 
trainData_1 <- titanic_imputed[trainIdx_1, ]
titanic_ranger_model_1 <- ranger(survived~., data = trainData_1, num.trees = 50, probability = TRUE) 
exp_1 <- explain(titanic_ranger_model_1, data = trainData_1) 
model_profile_1 <- model_profile(exp_1, type = "partial", grid = common_grid)

# 模型2
trainIdx_2 <- sample(nrow(titanic_imputed), 2/3 * nrow(titanic_imputed)) 
trainData_2 <- titanic_imputed[trainIdx_2, ]
titanic_ranger_model_2 <- ranger(survived~., data = trainData_2, num.trees = 50, probability = TRUE)
exp_2 <- explain(titanic_ranger_model_2, data = trainData_2)
model_profile_2 <- model_profile(exp_2, type = "partial", grid = common_grid)

# 合并所有profile数据(500个模型时,用list存储后执行bind_rows)
all_profiles <- bind_rows(
  model_profile_1$agr_profiles %>% mutate(model_id = 1),
  model_profile_2$agr_profiles %>% mutate(model_id = 2)
)

# 计算平均profile:按变量、x值分组,求预测值的均值
average_profile <- all_profiles %>%
  group_by(_vname_, _x_, _label_) %>%
  summarise(avg_yhat = mean(_yhat_), .groups = "drop")

# 转换为model_profile类对象,方便用DALEX的plot函数可视化
average_profile_obj <- list(
  agr_profiles = average_profile %>% rename(_yhat_ = avg_yhat),
  type = "partial",
  variable_type = "both",
  label = "Average Partial Dependence Profile"
)
class(average_profile_obj) <- c("model_profile", "explainer")

# 绘制平均profile
plot(average_profile_obj)

场景2:已生成500个model_profile(x值不一致)

如果已经有了所有模型的profile结果,可通过插值将每个模型的曲线对齐到统一的x网格,再计算均值。

实现步骤:

  1. 收集所有模型的profile数据
  2. 对每个变量生成覆盖全范围的统一x网格
  3. 对每个模型的曲线进行插值,得到统一x点上的预测值
  4. 计算所有模型在每个x点的平均预测值
library("ranger") 
library(DALEX)
library(dplyr)

# 假设你已有包含500个model_profile对象的列表:profiles_list
profiles_list <- list(model_profile_1, model_profile_2) # 示例,替换为你的500个对象

# 定义插值函数:针对单个变量,对齐所有模型的曲线并求平均
interpolate_and_average <- function(var_name, profiles_list) {
  # 收集该变量的所有模型的x和预测值
  var_data_list <- lapply(profiles_list, function(profile) {
    profile$agr_profiles %>% 
      filter(_vname_ == var_name) %>% 
      select(_x_, _yhat_)
  })
  
  # 生成统一的x网格:覆盖所有模型的x值范围,取100个均匀点
  all_x_values <- unlist(lapply(var_data_list, function(d) d$_x_))
  common_x <- seq(min(all_x_values), max(all_x_values), length.out = 100)
  
  # 对每个模型的曲线插值到common_x
  interpolated_y_list <- lapply(var_data_list, function(d) {
    approx(d$_x_, d$_yhat_, xout = common_x)$y
  })
  
  # 计算每个x点的平均预测值
  avg_yhat <- rowMeans(do.call(cbind, interpolated_y_list), na.rm = TRUE)
  
  # 返回结果数据框
  data.frame(
    _vname_ = var_name,
    _x_ = common_x,
    _yhat_ = avg_yhat,
    _label_ = "Average Partial Dependence Profile"
  )
}

# 获取所有变量名
all_var_names <- unique(profiles_list[[1]]$agr_profiles$_vname_)

# 对每个变量执行插值和平均计算
average_profile <- bind_rows(lapply(all_var_names, interpolate_and_average, profiles_list = profiles_list))

# 转换为model_profile对象并可视化
average_profile_obj <- list(
  agr_profiles = average_profile,
  type = "partial",
  variable_type = "both",
  label = "Average Partial Dependence Profile"
)
class(average_profile_obj) <- c("model_profile", "explainer")

plot(average_profile_obj)

内容的提问来源于stack exchange,提问作者Jason Connelly

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 15:26:10