如何用mlr3pipelines移除高比例NA值的列?我的尝试未生效
解决mlr3pipelines移除高NA占比列的问题
问题原因
你当前的代码中,通过参数转换(trafo)动态生成的selector函数无法正确捕获na_cutoff变量的值,导致筛选逻辑未按预期执行——本该被移除的Sepal.Width列(NA占比约33%,超过默认阈值0.2)没有被过滤掉。
修复方案1:修正作用域绑定问题
通过force()函数强制绑定na_cutoff的当前值,避免延迟解析导致的作用域问题:
library(mlr3) library(mlr3pipelines) task = tsk("iris") dt = task$data() dt[1:50, Sepal.Width := NA] task_ = as_task_classif(dt, target = "Species") graph = po("removeconstants", id = "removeconstants", ratio = 0.01) %>>% po("select", id = "drop_na_cols") ps = ParamSet$new(list(ParamDbl$new("na_cutoff", lower = 0, upper = 1, default = 0.2))) graph$param_set$add(ps) graph$param_set$trafo = function(x, param_set) { na_cutoff = x$na_cutoff # 强制绑定当前na_cutoff的值,避免作用域延迟解析 force(na_cutoff) x$drop_na_cols.selector = function(task) { fn = task$feature_names data = task$data(cols = fn) na_ratios = colMeans(is.na(data)) # 保留NA占比不超过阈值的列 fn[na_ratios <= na_cutoff] } x$na_cutoff = NULL x } train_res = graph$train(task_) # 验证结果:Sepal.Width已被移除 train_res$drop_na_cols.output$data()
修复方案2:自定义PipeOp(更直观可靠)
直接创建一个专门用于移除高NA占比列的PipeOp,逻辑清晰且避免作用域问题:
library(mlr3) library(mlr3pipelines) # 自定义移除高NA列的PipeOp PipeOpRemoveHighNA = R6::R6Class("PipeOpRemoveHighNA", inherit = PipeOpTaskPreproc, public = list( initialize = function(id = "remove_high_na", param_vals = list()) { ps = ParamSet$new(list( ParamDbl$new("na_cutoff", lower = 0, upper = 1, default = 0.2) )) super$initialize(id, param_set = ps, param_vals = param_vals) } ), private = list( .train_task = function(task) { na_cutoff = self$param_set$values$na_cutoff fn = task$feature_names na_ratios = colMeans(is.na(task$data(cols = fn))) self$state$keep_cols = fn[na_ratios <= na_cutoff] task$select(self$state$keep_cols) }, .predict_task = function(task) { # 预测阶段复用训练时确定的保留列,保证前后一致 task$select(self$state$keep_cols) } ) ) # 测试代码 task = tsk("iris") dt = task$data() dt[1:50, Sepal.Width := NA] task_ = as_task_classif(dt, target = "Species") # 构建管道 graph = po("removeconstants", ratio = 0.01) %>>% PipeOpRemoveHighNA$new(param_vals = list(na_cutoff = 0.2)) train_res = graph$train(task_) # 验证结果 train_res[[1]]$data()
内容的提问来源于stack exchange,提问作者Mislav Sagovac
相关产品推荐
相关产品推荐

