{"spec_id":"survival-kaplan-meier","library":"plotnine","language":"python","code":"\"\"\" anyplot.ai\nsurvival-kaplan-meier: Kaplan-Meier Survival Plot\nLibrary: plotnine 0.15.4 | Python 3.13.13\nQuality: 91/100 | Updated: 2026-05-11\n\"\"\"\n\nimport os\n\nimport numpy as np\nimport pandas as pd\nfrom plotnine import (\n    aes,\n    element_blank,\n    element_line,\n    element_rect,\n    element_text,\n    geom_point,\n    geom_ribbon,\n    geom_step,\n    geom_vline,\n    ggplot,\n    labs,\n    scale_color_manual,\n    scale_fill_manual,\n    scale_x_continuous,\n    scale_y_continuous,\n    theme,\n    theme_minimal,\n)\n\n\n# Theme tokens\nTHEME = os.getenv(\"ANYPLOT_THEME\", \"light\")\nPAGE_BG = \"#FAF8F1\" if THEME == \"light\" else \"#1A1A17\"\nELEVATED_BG = \"#FFFDF6\" if THEME == \"light\" else \"#242420\"\nINK = \"#1A1A17\" if THEME == \"light\" else \"#F0EFE8\"\nINK_SOFT = \"#4A4A44\" if THEME == \"light\" else \"#B8B7B0\"\n\n# Okabe-Ito palette\nIMPRINT = [\"#009E73\", \"#C475FD\"]\n\n# Data generation\nnp.random.seed(42)\n\n\n# Generate survival data for two treatment groups\ndef generate_times_and_events(n, hazard_rate):\n    times = np.random.exponential(1 / hazard_rate, n)\n    censor_times = np.random.uniform(0, np.percentile(times, 80), n)\n    censored = times > censor_times\n    observed_times = np.where(censored, censor_times, times)\n    events = (~censored).astype(int)\n    return observed_times, events\n\n\ntimes_a, events_a = generate_times_and_events(80, 0.02)\ntimes_b, events_b = generate_times_and_events(80, 0.035)\n\ndata = pd.concat(\n    [\n        pd.DataFrame({\"time\": times_a, \"event\": events_a, \"group\": \"Treatment A\"}),\n        pd.DataFrame({\"time\": times_b, \"event\": events_b, \"group\": \"Treatment B\"}),\n    ],\n    ignore_index=True,\n)\n\n\n# Kaplan-Meier estimation\ndef compute_km(df):\n    df = df.sort_values(\"time\").reset_index(drop=True)\n    n = len(df)\n    times = [0]\n    survival = [1.0]\n    ci_lower = [1.0]\n    ci_upper = [1.0]\n    var_sum = 0\n\n    at_risk = n\n    current_survival = 1.0\n    unique_times = df[df[\"event\"] == 1][\"time\"].unique()\n    unique_times.sort()\n\n    for t in unique_times:\n        at_risk = (df[\"time\"] >= t).sum()\n        events = ((df[\"time\"] == t) & (df[\"event\"] == 1)).sum()\n\n        if at_risk > 0:\n            current_survival *= (at_risk - events) / at_risk\n            if at_risk > events:\n                var_sum += events / (at_risk * (at_risk - events))\n\n            times.append(t)\n            survival.append(current_survival)\n\n            if current_survival > 0 and var_sum > 0:\n                se = current_survival * np.sqrt(var_sum)\n                z = 1.96\n                log_surv = np.log(current_survival)\n                log_se = se / current_survival\n                ci_lower.append(np.exp(log_surv - z * log_se))\n                ci_upper.append(np.exp(log_surv + z * log_se))\n            else:\n                ci_lower.append(current_survival)\n                ci_upper.append(current_survival)\n\n    max_time = df[\"time\"].max()\n    times.append(max_time)\n    survival.append(survival[-1])\n    ci_lower.append(ci_lower[-1])\n    ci_upper.append(ci_upper[-1])\n\n    return pd.DataFrame(\n        {\"time\": times, \"survival\": survival, \"ci_lower\": np.clip(ci_lower, 0, 1), \"ci_upper\": np.clip(ci_upper, 0, 1)}\n    )\n\n\nkm_a = compute_km(data[data[\"group\"] == \"Treatment A\"])\nkm_a[\"group\"] = \"Treatment A\"\n\nkm_b = compute_km(data[data[\"group\"] == \"Treatment B\"])\nkm_b[\"group\"] = \"Treatment B\"\n\nkm_data = pd.concat([km_a, km_b], ignore_index=True)\n\n# Get censored observations for tick marks\ncensored = data[data[\"event\"] == 0].copy()\ncensored_marks = []\nfor _, row in censored.iterrows():\n    group = row[\"group\"]\n    t = row[\"time\"]\n    km_group = km_a if group == \"Treatment A\" else km_b\n    surv = km_group[km_group[\"time\"] <= t][\"survival\"].iloc[-1]\n    censored_marks.append({\"time\": t, \"survival\": surv, \"group\": group})\n\ncensored_df = pd.DataFrame(censored_marks)\n\n\n# Compute median survival times for annotations\ndef get_median_survival(km_group):\n    surv_below_half = km_group[km_group[\"survival\"] <= 0.5]\n    if len(surv_below_half) > 0:\n        return surv_below_half.iloc[0][\"time\"]\n    return None\n\n\nmedian_a = get_median_survival(km_a)\nmedian_b = get_median_survival(km_b)\n\n# Plot\nanyplot_theme = theme(\n    plot_background=element_rect(fill=PAGE_BG, color=PAGE_BG),\n    panel_background=element_rect(fill=PAGE_BG, color=PAGE_BG),\n    panel_grid_major=element_line(color=INK, size=0.3, alpha=0.08),\n    panel_grid_minor=element_blank(),\n    axis_ticks_y=element_line(color=INK_SOFT, size=0.3),\n    axis_ticks_x=element_blank(),\n    panel_border=element_blank(),\n    axis_line_y=element_line(color=INK_SOFT, size=0.4),\n    axis_line_x=element_blank(),\n    axis_title=element_text(color=INK, size=20),\n    axis_text=element_text(color=INK_SOFT, size=16),\n    plot_title=element_text(color=INK, size=24, face=\"medium\"),\n    legend_background=element_rect(fill=ELEVATED_BG, color=INK_SOFT, size=0.4),\n    legend_text=element_text(color=INK_SOFT, size=16),\n    legend_title=element_text(color=INK, size=18),\n    legend_position=(0.70, 0.25),\n    legend_box_just=\"left\",\n    figure_size=(16, 9),\n)\n\nplot = (\n    ggplot()\n    + geom_ribbon(km_data, aes(x=\"time\", ymin=\"ci_lower\", ymax=\"ci_upper\", fill=\"group\"), alpha=0.15)\n    + geom_step(km_data, aes(x=\"time\", y=\"survival\", color=\"group\"), size=1.5)\n    + geom_point(censored_df, aes(x=\"time\", y=\"survival\", color=\"group\"), shape=\"|\", size=4)\n)\n\nif median_a is not None:\n    plot = plot + geom_vline(xintercept=median_a, linetype=\"dashed\", color=IMPRINT[0], alpha=0.4, size=0.8)\n\nif median_b is not None:\n    plot = plot + geom_vline(xintercept=median_b, linetype=\"dashed\", color=IMPRINT[1], alpha=0.4, size=0.8)\n\nplot = (\n    plot\n    + scale_color_manual(values=IMPRINT)\n    + scale_fill_manual(values=IMPRINT)\n    + scale_y_continuous(limits=(0, 1.05), breaks=[0, 0.25, 0.5, 0.75, 1.0], labels=[\"0%\", \"25%\", \"50%\", \"75%\", \"100%\"])\n    + scale_x_continuous(limits=(0, None))\n    + labs(\n        title=\"survival-kaplan-meier · plotnine · anyplot.ai\",\n        x=\"Time (months)\",\n        y=\"Survival Probability\",\n        color=\"Treatment Group\",\n        fill=\"Treatment Group\",\n    )\n    + theme_minimal()\n    + anyplot_theme\n)\n\n# Save\nplot.save(f\"plot-{THEME}.png\", dpi=300, verbose=False)\n"}