mlr3调参时Deepsurv与Loghaz在高删失率下的训练错误排查
问题:严重删失场景下mlr3生存分析模型的报错与解决方案
背景与问题描述
我正在研究严重删失场景下的机器学习性能,使用mlr3框架开展生存分析建模时遇到两类核心问题:
1. Deepsurv高删失率交叉验证失败
- 现象:50%删失率时,带参数调优的Deepsurv在5折交叉验证中有2折失败;90%删失率时5折全部失败,报错:
Error in check_prediction_data.PredictionDataSurv(pdata, train_task = task) : Assertion on 'pdata$crank' failed: Contains missing values (element 1).
- 特殊情况:90%删失率下改用3折交叉验证,带调参的Deepsurv可正常运行。
- 已尝试无效方案:参考mlr3 GitHub提示(极不平衡数据集导致部分重采样无事件样本),使用@mllg提出的任务实例化方案,未解决问题。
2. Survival Logistic-Hazard Learner概率验证错误
部分运行中出现如下报错:
Error in assert_surv_matrix(x) : Survival probabilities must be (non-strictly) decreasing and between [0, 1]
猜想与需求
- 怀疑调参过程未在实例化时遵循任务的
stratum角色。 - 寻求无需减少折数的解决方案(拒绝将5折改为3折的简单处理)。
- 可提供复现代码与数据集。
针对性解决方案
解决Deepsurv缺失crank值问题
(1)强制分层交叉验证
确保交叉验证按删失状态(事件/非事件)分层,避免某一折中全为删失样本。创建重采样策略时指定stratum参数:
resampling = rsmp("cv", folds = 5, stratum = "status") # "status"为事件状态列名 resampling$instantiate(task)
(2)过滤无效重采样折
预检查每个重采样折的样本分布,跳过无事件样本的折:
resampling$instantiate(task) # 筛选包含至少1个事件样本的折 valid_splits = Filter(function(split) { train_data = task$data(rows = split$train) sum(train_data$status == 1) > 0 }, resampling$splits) # 基于有效折构建自定义重采样对象 custom_resampling = rsmp("custom") custom_resampling$splits = valid_splits
(3)自定义Learner补全缺失crank
当模型无法生成有效crank值时,用样本平均生存时间作为兜底值:
LearnerSurvDeepsurvCustom = R6::R6Class("LearnerSurvDeepsurvCustom", inherit = LearnerSurvDeepsurv, public = list( predict = function(task) { pred = super$predict(task) # 补全缺失的crank值 if (any(is.na(pred$crank))) { mean_time = mean(task$truth()$time) pred$crank[is.na(pred$crank)] = mean_time } pred } ) ) # 使用自定义Learner learner = LearnerSurvDeepsurvCustom$new()
解决Logistic-Hazard概率递减问题
(1)添加正则化约束
调参时增大模型正则化强度,避免过拟合导致概率输出异常:
# 扩展调参空间,提高正则化参数范围 ps = ParamSet$new(list( ParamDbl$new("penalty", lower = 1e-3, upper = 1e1), # 其他调参参数... ))
(2)后处理修正生存概率
对模型输出的生存概率强制做递减修正:
LearnerSurvLogisticHazardCustom = R6::R6Class("LearnerSurvLogisticHazardCustom", inherit = LearnerSurvLogisticHazard, public = list( predict = function(task) { pred = super$predict(task) # 修正生存概率为递减序列 surv_matrix = pred$distr$survival corrected_surv = t(apply(surv_matrix, 1, cummin)) # 更新分布对象 pred$distr = distr6::VectorDistribution$new( distribution = "Survival", survival = corrected_surv, params = list(time = pred$distr$parameters$time) ) pred } ) ) # 使用自定义Learner learner = LearnerSurvLogisticHazardCustom$new()
(3)调参时加入概率有效性惩罚
自定义评估指标,对不符合递减要求的预测给予惩罚:
custom_cindex = msr("surv.cindex", id = "custom_cindex") custom_cindex$score = function(prediction, task, ...) { surv_matrix = prediction$distr$survival # 检查每行生存概率是否递减 is_valid = apply(surv_matrix, 1, function(x) all(x == cummin(x))) if (!all(is_valid)) { return(-Inf) # 惩罚无效预测 } super$score(prediction, task, ...) } # 调参时使用自定义指标 tuner = tnr("grid_search") tuning_instance = ti( task = task, learner = learner, resampling = resampling, measure = custom_cindex, search_space = ps, terminator = trm("evals", n_evals = 50) ) tuner$optimize(tuning_instance)
内容的提问来源于stack exchange,提问作者A. Suliman
相关产品推荐
相关产品推荐

