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
相关产品推荐
相关产品推荐

