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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.20 01:22:06