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

如何查看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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 23:10:18