如何通过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
相关产品推荐
相关产品推荐

