如何用mlr3实现按月滚动交叉验证?现有方案遇问题
用mlr3实现按月滚动交叉验证(Rolling CV)的方案
需求说明
需要实现按月滚动交叉验证:用连续6个月的数据作为训练集,1个月数据作为测试集,循环滚动完成所有可用月份的验证。
示例数据集
DT <- structure(list(Not_FLS_positive = c(0.408197129345391, 0.765784452003651, 0.44694266987472, 0.261843524433751, 0.823612378660914, 0.463701982908819, 0.50286235791919, 0.202937028125778, 0.728864183190907, 0.396498796980005, 0.0645482452501452, 0.386210901850162, 0.518874968887414, 0.748527337592301, 0.453414087778976, 0.758566332033519, 0.544926574296856, 0.758151497552477, 0.641583008379657, 0.15000414834481, 0.271384717497718, 0.516634862689787, 0.379988384634531, 0.220277109433336, 0.368373019165353, 0.367294449514644, 0.924583091346553, 0.702895544677674, 0.560192483199204, 0.61212976022567, 0.0189164523355181, 0.308139052518045), Not_FLS_negative = c(0.690284576453995, 0.406288890732598, 0.965402804281092, 0.981830249730358, 0.750850410686136, 0.884676014270306, 0.978760474570646, 0.846013440637186, 0.319754417987223, 0.70256367709284, 0.0308636853895296, 0.247905085870738, 0.886999087364142, 0.28017920849581, 0.697253795735502, 0.720069692192815, 0.838131585497387, 0.967559943582511, 0.755745457562433, 0.97593960009956, 0.886833153571725, 0.587156724466938, 0.959097320169252, 0.0548411183937609, 0.957769849829918, 0.479382726292209, 0.626897867750767, 0.772670704388949, 0.9822450842114, 0.736829005226914, 0.420642163776653, 0.723886169418402), bin_aroundzero_ret_excess_stand_22 = structure(c(2L, 1L, 3L, 1L, 1L, 3L, 1L, 1L, 2L, 2L, 2L, 1L, 3L, 1L, 2L, 2L, 1L, 1L, 1L, 3L, 2L, 1L, 3L, 2L, 2L, 2L, 3L, 2L, 1L, 2L, 3L, 2L), levels = c("0", "1", "-1"), class = "factor"), monthid = c("20141", "20141", "20141", "20141", "20141", "20141", "20141", "20141", "20141", "20141", "20142", "20142", "20142", "20142", "20142", "20142", "20142", "20142", "20142", "20142", "20142", "20143", "20143", "20143", "20143", "20143", "20143", "20143", "20143", "20143", "20143", "20143")), row.names = c(NA, -32L), class = c("data.table", "data.frame"))
原尝试的错误原因
使用mlr3temporal的forecast_cv时出现以下错误:
Error in max(ids) - self$param_set$values$horizon : non-numeric argument to binary operator
这是因为forecast_cv要求分组变量必须是数值型,但原数据中的monthid是字符型(如"20141"),导致内部数值计算失败。
解决方案
方案1:修正forecast_cv的使用(推荐)
将monthid转为数值型,再使用官方的滚动窗口重采样实现:
# 转换monthid为整数型 DT$monthid <- as.integer(DT$monthid) # 创建分类任务并设置monthid为分组角色 task <- as_task_classif(DT, id = "aroundzero", target = "bin_aroundzero_ret_excess_stand_22") task$set_col_roles("monthid", "group") # 初始化滚动窗口重采样:固定6个月窗口,每次测试1个月 resampling <- rsmp("forecast_cv", fixed_window = TRUE, horizon = 1L, window_size = 6) resampling$instantiate(task) # 验证每一轮的训练/测试月份 for (i in seq_len(resampling$iters)) { train_months <- unique(DT$monthid[resampling$train_set(i)]) test_months <- unique(DT$monthid[resampling$test_set(i)]) cat(sprintf("第%d轮: 训练月份=%s, 测试月份=%s\n", i, paste(train_months, collapse = ","), paste(test_months, collapse = ","))) }
方案2:自定义按月分组的重采样
如果不想修改monthid的类型,可以手动生成按月划分的训练/测试集索引:
# 获取按顺序排列的唯一月份 unique_months <- unique(DT$monthid) n_months <- length(unique_months) # 生成训练/测试的月份组合:6个月训练,1个月测试 train_month_sets <- list() test_month_sets <- list() for (i in seq(from = 7, to = n_months)) { train_month_sets[[i-6]] <- unique_months[(i-6):(i-1)] test_month_sets[[i-6]] <- unique_months[i] } # 将月份映射为样本索引 train_sets <- lapply(train_month_sets, function(months) which(DT$monthid %in% months)) test_sets <- lapply(test_month_sets, function(month) which(DT$monthid == month)) # 初始化自定义重采样 custom_resampling <- rsmp("custom") custom_resampling$instantiate(task, train_sets, test_sets) # 查看第一轮的训练/测试样本索引 custom_resampling$train_set(1) custom_resampling$test_set(1)
方案选择
- 如果
monthid可以转换为数值型,优先使用方案1,利用官方实现的稳定性和便捷性; - 如果
monthid是无法转为数值的格式(如"2014-01"),则使用方案2,自定义分组逻辑更灵活。
内容的提问来源于stack exchange,提问作者Mislav Sagovac
相关产品推荐
相关产品推荐

