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

如何通过XGBoost AFT实现获取生存与风险评分?

问题描述

我用XGBoost的AFT生存分析模型完成训练后,调用xgb.predict(dtest)只能得到每个样本的事件发生时间预测值。想搞清楚:

  • 能不能获取各时间点对应的概率评分?
  • 如何得到样本的生存值与风险值?
    之前用lifelines包可以输出各时间窗口的概率矩阵,想知道XGBoost有没有类似的实现方式。

附训练代码:

dtrain = xgb.DMatrix(X)

# mlmodel.E_train是布尔型事件指示器
# mlmodel.T_train是事件观测或截尾时的耗时
y_lower_bound = np.where(mlmodel.E_train == 1, mlmodel.T_train, 0)
y_upper_bound = np.where(mlmodel.E_train == 1, mlmodel.T_train, np.inf)

dtrain.set_float_info('label_lower_bound', y_lower_bound)
dtrain.set_float_info('label_upper_bound', y_upper_bound)

params = {
    'objective': 'survival:aft',
    'eval_metric': 'aft-nloglik',
    'aft_loss_distribution': 'normal',
    'aft_loss_distribution_scale': 1.20,
    'tree_method': 'hist', 
    'learning_rate': 0.05, 
    'max_depth': 2
}

bst = xgb.train(params, dtrain, num_boost_round=5, evals=[(dtrain, 'train')])

mlmodel.model = bst

dtest = xgb.DMatrix(self.X_test)
mlmodel.y_pred_hazard = pd.DataFrame(mlmodel.model.predict(dtest))
解决方案

XGBoost的AFT模型默认只输出事件发生时间的条件均值(对应normal分布)或中位数(对应logistic分布),但可以基于它的分布假设手动推导出生存函数、风险函数以及各时间点的概率值,具体步骤如下:

1. 明确AFT模型的分布逻辑

你用的是aft_loss_distribution: normal,AFT模型的核心假设是事件时间的对数服从正态分布:

  • $ln(T) = \eta + \epsilon$,其中$\eta$是模型学习的线性预测器,$\epsilon \sim N(0, \sigma^2)$,$\sigma$就是你设置的aft_loss_distribution_scale
  • predict方法返回的是事件时间的条件均值:$E[T] = exp(\eta + \sigma^2/2)$,所以需要先反推得到$\eta$

2. 计算各时间点的生存概率

生存函数$S(t)$表示“到时间$t$时事件仍未发生的概率”,公式为:
$S(t) = P(T > t) = 1 - \Phi\left( \frac{ln(t) - \eta}{\sigma} \right)$
其中$\Phi$是标准正态分布的累积分布函数(CDF)。

3. 计算各时间点的风险值

风险函数$h(t)$表示“在时间$t$时事件发生的瞬时概率”,公式为:
$h(t) = \frac{f(t)}{S(t)} = \frac{\phi\left( \frac{ln(t) - \eta}{\sigma} \right)}{t \cdot \sigma \cdot (1 - \Phi\left( \frac{ln(t) - \eta}{\sigma} \right))}$
其中$\phi$是标准正态分布的概率密度函数(PDF)。

4. 代码实现

基于你的现有代码,添加以下逻辑即可得到生存/风险矩阵:

import numpy as np
import pandas as pd
from scipy.stats import norm

# 获取模型预测值与尺度参数
y_pred = mlmodel.model.predict(dtest)
sigma = params['aft_loss_distribution_scale']

# 反推线性预测器η
eta = np.log(y_pred) - 0.5 * sigma ** 2

# 定义需要计算的时间点(可自定义,这里用训练集中的唯一观测时间)
time_points = np.sort(np.unique(mlmodel.T_train[mlmodel.T_train > 0]))

# 计算每个样本在各时间点的生存概率
survival_data = []
for t in time_points:
    z = (np.log(t) - eta) / sigma
    survival_prob = 1 - norm.cdf(z)
    survival_data.append(survival_prob)
survival_df = pd.DataFrame(np.array(survival_data).T, columns=[f"t_{round(t,2)}" for t in time_points])

# 计算每个样本在各时间点的风险值
hazard_data = []
for t in time_points:
    z = (np.log(t) - eta) / sigma
    pdf_val = norm.pdf(z)
    cdf_val = norm.cdf(z)
    hazard_val = pdf_val / (t * sigma * (1 - cdf_val))
    hazard_data.append(hazard_val)
hazard_df = pd.DataFrame(np.array(hazard_data).T, columns=[f"t_{round(t,2)}" for t in time_points])

# 最终结果:
# survival_df:每行是一个样本,每列对应一个时间点的生存概率
# hazard_df:每行是一个样本,每列对应一个时间点的风险值

5. 额外说明

  • 如果换用其他分布(比如logistic、extreme),只需要替换scipy.stats中的对应分布(如logistic、gumbel_r)即可,公式逻辑一致。
  • XGBoost的AFT模型没有内置输出生存/风险矩阵的API,这和lifelines不同——lifelines的Cox模型等是直接基于比例风险假设输出这类结果,而AFT模型需要手动基于分布推导。
  • 训练时的截尾数据已经通过label_lower_bound和label_upper_bound纳入模型学习,预测阶段不需要额外处理。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.22 12:15:42