如何为DataFrame中所有object类型列绘制Kaplan-Meier生存曲线子图
修正Kaplan-Meier生存曲线子图绘制问题
问题说明
你当前的代码会将所有object列的生存曲线重复绘制到每个子图中,核心原因是没有将每个目标列与对应的子图轴(ax)建立一一对应的绑定关系,而是循环所有轴进行绘图。参考你提供的柱状图逻辑,我们可以通过以下步骤修正:
修正方向
- 先筛选出DataFrame中所有
object类型的列,避免遍历无关列 - 根据筛选出的列数量配置子图布局(确保子图数量与目标列数量一致)
- 将筛选出的列与子图轴一一配对,在对应轴上仅绘制该列的所有分组生存曲线
- 为每个子图添加标题、图例和坐标轴标签,提升可读性
完整修正代码
首先补全示例数据中缺失的gender变量定义,再编写修正后的生存曲线代码:
import numpy as np import random import pandas as pd import matplotlib.pyplot as plt from lifelines import KaplanMeierFitter # 补全示例数据 duration = np.random.exponential(scale=5, size=100).round(1) boolean = [bool(random.randint(0, 1)) for _ in range(len(duration))] group = np.random.choice(["A", "B", "C", "D"], size=len(duration)) house = np.random.choice(["Big", "Small"], p=[0.7, 0.3], size=len(duration)) provider = np.random.choice(["2Degrees", "Skinny", "Vodafone", "Spark"], p=[0.25]*4, size=len(duration)) gender = np.random.choice(["Male", "Female"], size=len(duration)) # 补全缺失的gender变量 df = pd.DataFrame({ "Duration": duration, "Boolean": boolean, "Group": group, "Gender": gender, "Provider": provider }) # 修正后的生存曲线绘制代码 kmf = KaplanMeierFitter() # 筛选所有object类型的列 cat_cols = [col for col in df.columns if df[col].dtype == object] # 根据列数设置子图布局 fig, axes = plt.subplots(nrows=1, ncols=len(cat_cols), figsize=(15, 5)) # 将列与对应轴一一配对绘制 for col, ax in zip(cat_cols, axes.flatten()): # 遍历当前列的每个唯一值,绘制KM曲线 for value in df[col].unique(): mask = df[col] == value kmf.fit( durations=df["Duration"][mask], event_observed=df["Boolean"][mask], label=value ) kmf.plot_survival_function(ax=ax, ci_show=False) # 设置子图标题和坐标轴标签 ax.set_title(f"Kaplan-Meier Curve: {col}") ax.set_xlabel("Duration") ax.set_ylabel("Survival Probability") ax.legend(title=col) plt.tight_layout() plt.show()
代码说明
- 筛选目标列:用列表推导式直接筛选出所有
object类型的列,存储到cat_cols中 - 子图布局:根据
cat_cols的长度设置子图数量,保证每个列对应一个子图 - 一一配对绘图:通过
zip(cat_cols, axes.flatten())将每个列和对应的轴绑定,仅在当前轴上绘制该列的分组曲线 - 美化图表:添加标题、坐标轴标签和图例,让每个子图的内容更清晰
内容的提问来源于stack exchange,提问作者JoMcGee
相关产品推荐
相关产品推荐

