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

mlr3中surv.xgboost.cox如何获取生存预测时间并与真实数据对比?

解决mlr3中surv.xgboost.cox预测生存时间的问题

问题说明

使用surv.xgboost.cox学习器时,在构造学习器时传入type="regression"会报错,因为该参数并非此学习器的有效参数。Cox比例风险模型本身不直接输出生存时间,需通过正确方式获取生存时间预测值与真实值对比。

修正步骤及代码

  1. 基础准备:数据处理与划分
library(mlr3proba)
library(mlr3extralearners)
library(mlr3pipelines)
library(mlr3verse)
library(survival)

# 创建生存分析任务
task = as_task_surv(x = veteran, time = 'time', event = 'status')
# 编码分类变量
poe = po('encode')
task = poe$train(list(task))[[1]]
# 划分训练集与测试集(修正原代码中变量名错误)
set.seed(42)
part = partition(task, ratio = 0.8)
  1. 模型训练与生存分布预测
    surv.xgboost.cox支持的预测类型为"risk"(默认,输出风险分数)和"distr"(输出生存分布)。要获取生存时间点估计,需在predict方法中指定type="distr":
# 初始化学习器,无需设置type参数
learner = lrn("surv.xgboost.cox")
# 训练模型
learner$train(task, part$train)
# 预测生存分布
pred_xgb = learner$predict(task, part$test, type = "distr")
  1. 提取预测生存时间并与真实值对比
    从生存分布中提取中位数(或其他分位数)作为预测生存时间,与真实数据对比:
# 提取中位数生存时间作为预测值
pred_time = pred_xgb$distr$quantile(0.5)
# 获取真实生存时间与事件状态
truth_time = pred_xgb$truth[, "time"]
truth_event = pred_xgb$truth[, "status"]

# 查看前3条数据的对比结果
compare_df = data.frame(
  预测生存时间 = pred_time[1:3],
  真实生存时间 = truth_time[1:3],
  事件状态 = truth_event[1:3]
)
print(compare_df)

注意事项

  • Cox模型的核心是风险比例估计,生存时间预测是基于基线生存函数推导的间接估计值。
  • 真实数据中存在删失样本(事件状态=0表示未发生事件),单纯对比点生存时间需谨慎,建议结合C-index、Brier分数等生存分析专用评估指标综合判断模型性能。

内容的提问来源于stack exchange,提问作者Faiza

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 15:13:15