如何查看mlr3中classbalancing pipeline处理后的数据集属性?
查看mlr3 Pipeline中classbalancing处理后的数据信息
要获取交叉验证中每个fold下采样后的行数、行索引,可通过以下两种方式实现:
方法1:添加PipeOpDebug捕获预处理后的数据
在构建Graph时插入PipeOpDebug算子,它会保存经过前序步骤处理后的数据,后续可从训练后的GraphLearner实例中提取:
library(mlr3) library(mlr3pipelines) library(mlr3learners) # 构建包含debug算子的管道 graph = po("classbalancing", id = "bal", adjust = "downsampling", reference = "minor", ratio = 5) %>>% po("encode") %>>% po("scale") %>>% po("debug", id = "debug_balanced") %>>% # 插入debug算子保存中间数据 po("learner", lrn("classif.cv_glmnet")) glrn = GraphLearner$new(graph) # 执行交叉验证(替换成你的任务) res = resample(tsk("your_task"), glrn, rsmp("cv", folds = 5)) # 提取第1个fold的平衡后数据 balanced_data = res$learners[[1]]$graph$pipeops$debug_balanced$output[[1]] cat("平衡后行数:", nrow(balanced_data), "\n") cat("平衡后行索引:", paste(balanced_data$row_ids, collapse = ", "), "\n") # 遍历所有fold查看 for (i in seq_along(res$learners)) { balanced_data = res$learners[[i]]$graph$pipeops$debug_balanced$output[[1]] cat(sprintf("Fold %d: 行数=%d, 行索引=%s\n", i, nrow(balanced_data), paste(balanced_data$row_ids, collapse = ", "))) }
方法2:用回调函数记录训练过程数据
通过自定义回调,在每个fold训练开始时执行预处理并记录数据信息,无需修改原管道结构:
# 定义回调函数 cb_record_balance = CallbackResample$new( id = "record_balanced_info", on_train_begin = function(ctx) { # 获取当前fold的训练子集 train_task = ctx$task$clone()$filter(ctx$row_ids$train) # 单独执行classbalancing及后续预处理步骤 preprocessed = graph$pipeops$bal$train(list(train_task))[[1]] preprocessed = graph$pipeops$encode$train(list(preprocessed))[[1]] preprocessed = graph$pipeops$scale$train(list(preprocessed))[[1]] # 记录信息 ctx$log$info(sprintf("Fold %d: 平衡后行数=%d,行索引=%s", ctx$iteration, nrow(preprocessed), paste(preprocessed$row_ids, collapse = ", "))) } ) # 带回调执行resample(替换成你的任务) res = resample(tsk("your_task"), glrn, rsmp("cv", folds = 5), callbacks = list(cb_record_balance)) # 查看日志中的记录 print(res$log)
注意事项
classbalancing的ratio参数表示多数类样本数/少数类样本数,设为5即多数类会被下采样到少数类数量的5倍。- 两种方法中,
balanced_data$row_ids对应原始任务的行索引,方便关联原始数据。
内容的提问来源于stack exchange,提问作者Jason Connelly
相关产品推荐
相关产品推荐

