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

R语言机器学习:如何在MLR3中使用MLR包的生存过滤器

可行实现方案

目前有两种成熟路径可以实现旧版MLR生存过滤器和MLR3生态的适配,无需修改现有过滤逻辑即可接入benchmark_grid()流程:


方案1:将MLR侧的融合管道封装为MLR3自定义生存学习器

如果你已经在MLR中完成了「生存过滤器+基础学习器」的融合管道,可以通过继承LearnerSurv基类的方式,将整个MLR管道的训练、预测逻辑封装为MLR3可识别的学习器对象,封装完成后即可直接传入MLR3的基准测试流程。
简化实现示例:

library(mlr3)
library(mlr3proba)
library(mlr)

# 先在MLR侧构建带生存过滤器的学习器
mlr_learner_with_filter = makeFilterWrapper(
  learner = makeLearner("surv.coxph"),
  fw.method = "coxph", # 可替换为任意MLR支持的生存过滤方法
  fw.abs = 10 # 保留特征数,可根据需求调整
)

# 自定义MLR3生存学习器,对接MLR管道
LearnerSurvMLRWrapper = R6::R6Class(
  "LearnerSurvMLRWrapper",
  inherit = LearnerSurv,
  public = list(
    initialize = function() {
      super$initialize(
        id = "surv.mlr_wrapper",
        feature_types = c("integer", "numeric", "factor"),
        predict_types = c("crank", "lp"),
        packages = "mlr"
      )
    }
  ),
  private = list(
    .train = function(task) {
      # 将MLR3生存任务转换为MLR生存任务
      mlr_task = mlr::makeSurvTask(
        data = task$data(),
        target = task$target_names,
        time = task$target_names[1],
        event = task$target_names[2]
      )
      # 训练MLR侧带过滤器的学习器
      model = mlr::train(mlr_learner_with_filter, mlr_task)
      return(model)
    },
    .predict = function(task) {
      pred = predict(self$model, newdata = task$data())
      # 将MLR预测结果转换为MLR3要求的格式
      return(list(lp = pred$data$response, crank = pred$data$response))
    }
  )
)

# 生成可直接用于MLR3流程的学习器
mlr3_learner = LearnerSurvMLRWrapper$new()

方案2:移植MLR生存过滤器为MLR3原生Filter对象

如果你希望直接在MLR3的特征选择框架中调用生存过滤逻辑,可以继承mlr3filters::Filter基类,把MLR过滤器的计算逻辑移植为MLR3原生过滤器,后续可以直接和MLR3的FilterWrapper、管道操作搭配使用,适配所有MLR3生存学习器。
简化实现示例:

library(mlr3filters)
library(survival)

# 自定义MLR3生存过滤器,以CoxPH单变量过滤为例
FilterSurvCoxPH = R6::R6Class(
  "FilterSurvCoxPH",
  inherit = Filter,
  public = list(
    initialize = function() {
      super$initialize(
        id = "surv.coxph",
        task_types = "surv",
        feature_types = c("integer", "numeric"),
        packages = "survival"
      )
    }
  ),
  private = list(
    .calculate = function(task, nfeat) {
      # 完全复用MLR中coxph过滤器的计算逻辑
      data = task$data()
      time_col = task$target_names[1]
      event_col = task$target_names[2]
      # 逐个计算特征和生存终点的关联得分
      scores = apply(task$data(cols = task$feature_names), 2, function(x) {
        mod = coxph(Surv(get(time_col), get(event_col)) ~ x, data = data)
        return(abs(summary(mod)$coefficients[,"z"]))
      })
      return(scores)
    }
  )
)

# 实例化后即可在MLR3中直接调用
coxph_surv_filter = FilterSurvCoxPH$new()

内容的提问来源于stack exchange,提问作者Mary B

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 17:54:03