如何从scikit-survival的CoxPHSurvivalAnalysis获取概率密度函数
获取并绘制Cox模型的概率密度函数$f(t)$
由于sksurv.linear_model.CoxPHSurvivalAnalysis没有直接输出概率密度函数的方法,但我们可以利用生存函数或累积风险函数的数学关系推导得到:
核心数学关系
概率密度函数$f(t)$与生存函数$S(t)$的关系为:
$$f(t) = -\frac{dS(t)}{dt}$$
对于sksurv返回的阶梯形式生存函数(StepFunction对象),我们可以通过计算相邻时间点的生存函数值差异,得到对应时间点的概率密度(或离散概率质量)。
实现步骤与代码修改
以下基于你的示例代码,添加获取和绘制概率密度函数的逻辑:
import pandas as pd import numpy as np import matplotlib.pyplot as plt from sksurv.linear_model import CoxPHSurvivalAnalysis data = np.random.randint(5,30,size=10) X_train = pd.DataFrame(data, columns=['covariate']) # 无删失场景:status全为True y_train = np.array([(True, t) for t in np.random.randint(0,100,size=10)/100], dtype=[('status',bool),('target',float)]) estimator = CoxPHSurvivalAnalysis() estimator.fit(X_train,y_train) X_test = pd.DataFrame({'covariate':[12,2]}) chf = estimator.predict_cumulative_hazard_function(X_test) survival_functions = estimator.predict_survival_function(X_test) # 绘制累积风险、生存函数、概率密度函数 fig, ax = plt.subplots(1,3, figsize=(15,5)) for fn_h, fn_s in zip(chf, survival_functions): # 绘制累积风险函数 ax[0].step(fn_h.x, fn_h(fn_h.x), where='post') # 绘制生存函数 ax[1].step(fn_s.x, fn_s(fn_s.x), where='post') # 计算概率密度函数 t_points = fn_s.x s_values = fn_s(t_points) # 初始生存值为1(t<第一个时间点时S(t)=1) f_values = np.zeros_like(s_values) f_values[0] = 1 - s_values[0] # 后续每个时间点的概率密度为前一个生存值减当前生存值 f_values[1:] = s_values[:-1] - s_values[1:] # 绘制概率密度函数 ax[2].step(t_points, f_values, where='post') ax[0].set_title('Cumulative Hazard Functions') ax[1].set_title('Survival Functions') ax[2].set_title('Probability Density Functions') plt.tight_layout() plt.show()
补充说明
- 由于sksurv返回的生存函数是阶梯函数,概率密度表现为每个事件时间点上的跳跃值(对应离散时间下的概率质量)
- 无删失场景下,所有概率密度的和应为1,可用于验证结果正确性
- 若需要连续形式的概率密度,可对生存函数使用
numpy.gradient进行数值微分,但需注意阶梯函数的微分特性
内容的提问来源于stack exchange,提问作者kevinsmith
相关产品推荐
相关产品推荐

